CoolFace
Apppublic

osky9/Land_Cover_classification

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
app.py361 linesDownload Raw Back to root
1import streamlit as st2import matplotlib3matplotlib.use("Agg")4 5st.set_page_config(6    page_title="Land Cover Classification",7    page_icon="satellite",8    layout="wide"9)10 11st.title("Land Cover Classification")12st.markdown("**Remote Sensing - MODIS Lake Powell - Machine Learning - Explainable AI**")13st.markdown("---")14 15try:16    import xgboost17    HAS_XGB = True18except ImportError:19    HAS_XGB = False20 21try:22    import lightgbm23    HAS_LGB = True24except ImportError:25    HAS_LGB = False26 27try:28    import shap29    HAS_SHAP = True30except ImportError:31    HAS_SHAP = False32 33st.sidebar.header("Configuration")34model_options = ["Random Forest"]35if HAS_XGB:36    model_options.append("XGBoost")37if HAS_LGB:38    model_options.append("LightGBM")39 40selected_model = st.sidebar.selectbox("Model", model_options)41test_size = st.sidebar.slider("Validation split", 0.10, 0.30, 0.15, 0.05)42if HAS_SHAP:43    run_xai = st.sidebar.checkbox("Run SHAP (slower)", value=False)44else:45    run_xai = False46    st.sidebar.info("SHAP not installed")47 48st.sidebar.markdown("---")49load_btn = st.sidebar.button("Load Dataset", use_container_width=True)50train_btn = st.sidebar.button("Train Model", use_container_width=True, type="primary")51 52if "data_loaded" not in st.session_state:53    st.session_state["data_loaded"] = False54if "model_trained" not in st.session_state:55    st.session_state["model_trained"] = False56 57if load_btn:58    with st.spinner("Downloading dataset..."):59        try:60            import pandas as pd61            BASE = (62                "https://huggingface.co/datasets/"63                "nasa-cisto-data-science-group/"64                "modis-lake-powell-toy-dataset/resolve/main/"65            )66            df_train = pd.read_csv(BASE + "train.csv")67            df_test = pd.read_csv(BASE + "test.csv")68            st.session_state["df_train"] = df_train69            st.session_state["df_test"] = df_test70            st.session_state["data_loaded"] = True71            st.success(f"Loaded! Train: {df_train.shape}, Test: {df_test.shape}")72        except Exception as e:73            st.error(f"Error: {e}")74 75tab1, tab2, tab3, tab4 = st.tabs(["EDA", "Features", "Model", "XAI"])76 77with tab1:78    if not st.session_state["data_loaded"]:79        st.info("Click Load Dataset in the sidebar to begin.")80    else:81        import numpy as np82        import matplotlib.pyplot as plt83        import seaborn as sns84        df_train = st.session_state["df_train"]85        df_test = st.session_state["df_test"]86        PALETTE = ["#64ffda","#f07178","#c3e88d","#82aaff","#ffcb6b","#ff5572"]87        KNOWN = ("label","class","target","landcover","land_cover","lc","category","type","cover")88        candidates = [c for c in df_train.columns if c.lower() in KNOWN]89        if not candidates:90            candidates = df_train.select_dtypes(exclude=np.number).columns.tolist()91        if not candidates:92            candidates = [c for c in df_train.columns if df_train[c].nunique() < 50]93        TARGET = candidates[0] if candidates else df_train.columns[-1]94        st.session_state["TARGET"] = TARGET95        num_features = df_train[[c for c in df_train.columns if c != TARGET]].select_dtypes(include=np.number).columns.tolist()96        st.session_state["num_features"] = num_features97        st.sidebar.success(f"Target: {TARGET}")98        st.sidebar.info(f"Classes: {df_train[TARGET].nunique()}")99        st.sidebar.info(f"Features: {len(num_features)}")100        c1, c2, c3, c4 = st.columns(4)101        c1.metric("Train rows", f"{len(df_train):,}")102        c2.metric("Test rows", f"{len(df_test):,}")103        c3.metric("Features", len(num_features))104        c4.metric("Classes", df_train[TARGET].nunique())105        st.dataframe(df_train.head(10), use_container_width=True)106        st.markdown("---")107        st.subheader("EDA 1 - Class Distribution")108        cc = df_train[TARGET].value_counts()109        fig, axes = plt.subplots(1, 2, figsize=(12, 4))110        axes[0].bar(cc.index.astype(str), cc.values, color=PALETTE[:len(cc)])111        axes[0].set_title("Count")112        axes[0].grid(axis="y", alpha=0.3)113        axes[1].pie(cc.values, labels=cc.index.astype(str), autopct="%1.1f%%", colors=PALETTE[:len(cc)])114        axes[1].set_title("Proportion")115        plt.tight_layout()116        st.pyplot(fig)117        plt.close()118        st.subheader("EDA 2 - Feature Distributions")119        n = min(len(num_features), 12)120        ncols = 4121        nrows = -(-n // ncols)122        fig, axes = plt.subplots(nrows, ncols, figsize=(16, nrows * 3))123        axes = axes.flatten()124        for i, col in enumerate(num_features[:n]):125            axes[i].hist(df_train[col].dropna(), bins=40, color=PALETTE[i % len(PALETTE)], alpha=0.85)126            axes[i].set_title(col, fontsize=9)127            axes[i].grid(axis="y", alpha=0.3)128        for j in range(i + 1, len(axes)):129            axes[j].set_visible(False)130        plt.tight_layout()131        st.pyplot(fig)132        plt.close()133        st.subheader("EDA 3 - Correlation Heatmap")134        corr = df_train[num_features].corr()135        mask = np.triu(np.ones_like(corr, dtype=bool))136        fig, ax = plt.subplots(figsize=(max(8, len(num_features) * 0.65), max(6, len(num_features) * 0.55)))137        sns.heatmap(corr, mask=mask, cmap="coolwarm", center=0, annot=(len(num_features) <= 15), fmt=".2f", linewidths=0.3, ax=ax)138        ax.set_title("Feature Correlation")139        plt.tight_layout()140        st.pyplot(fig)141        plt.close()142        st.subheader("EDA 4 - Per-Class Box Plots")143        top_feats = num_features[:min(6, len(num_features))]144        class_labels = df_train[TARGET].unique()145        fig, axes = plt.subplots(2, 3, figsize=(16, 8))146        axes = axes.flatten()147        for i, feat in enumerate(top_feats):148            groups = [df_train.loc[df_train[TARGET] == cl, feat].dropna().values for cl in class_labels]149            bp = axes[i].boxplot(groups, patch_artist=True, medianprops=dict(color="white", linewidth=1.5))150            for patch, color in zip(bp["boxes"], PALETTE):151                patch.set_facecolor(color)152                patch.set_alpha(0.8)153            axes[i].set_xticklabels([str(c) for c in class_labels], fontsize=9)154            axes[i].set_title(feat, fontsize=10)155            axes[i].grid(axis="y", alpha=0.3)156        for j in range(i + 1, len(axes)):157            axes[j].set_visible(False)158        plt.tight_layout()159        st.pyplot(fig)160        plt.close()161        st.subheader("EDA 5 - Descriptive Statistics")162        st.dataframe(df_train[num_features].describe().round(3), use_container_width=True)163 164with tab2:165    if not st.session_state["data_loaded"]:166        st.info("Load the dataset first.")167    else:168        import numpy as np169        df_train = st.session_state["df_train"]170        df_test = st.session_state["df_test"]171        TARGET = st.session_state["TARGET"]172        num_features = st.session_state["num_features"]173        st.subheader("Feature Engineering")174        def _eng(df, feats):175            df = df.copy()176            band_map = {}177            for col in feats:178                lc = col.lower()179                if any(k in lc for k in ("sur_refl_b01", "_red", "_b1")):180                    band_map["red"] = col181                if any(k in lc for k in ("sur_refl_b02", "_nir", "_b2")):182                    band_map["nir"] = col183                if any(k in lc for k in ("sur_refl_b06", "swir", "_b6")):184                    band_map["swir"] = col185                if any(k in lc for k in ("sur_refl_b04", "green", "_b4")):186                    band_map["green"] = col187            created = []188            if "nir" in band_map and "red" in band_map:189                nir = df[band_map["nir"]]190                red = df[band_map["red"]]191                df["NDVI"] = (nir - red) / (nir + red + 1e-8)192                created.append("NDVI")193            if "nir" in band_map and "swir" in band_map:194                nir = df[band_map["nir"]]195                swir = df[band_map["swir"]]196                df["NDWI"] = (nir - swir) / (nir + swir + 1e-8)197                created.append("NDWI")198            if "green" in band_map and "nir" in band_map:199                green = df[band_map["green"]]200                nir = df[band_map["nir"]]201                df["MNDWI"] = (green - nir) / (green + nir + 1e-8)202                created.append("MNDWI")203            df["feat_mean"] = df[feats].mean(axis=1)204            df["feat_std"] = df[feats].std(axis=1)205            df["feat_range"] = df[feats].max(axis=1) - df[feats].min(axis=1)206            created += ["feat_mean", "feat_std", "feat_range"]207            return df, created208        df_train_eng, created = _eng(df_train, num_features)209        df_test_eng, _ = _eng(df_test, num_features)210        num_features_eng = (df_train_eng[[c for c in df_train_eng.columns if c != TARGET]].select_dtypes(include=np.number).columns.tolist())211        st.session_state.update({212            "df_train_eng": df_train_eng,213            "df_test_eng": df_test_eng,214            "num_features_eng": num_features_eng215        })216        c1, c2 = st.columns(2)217        c1.metric("Original features", len(num_features))218        c2.metric("After engineering", len(num_features_eng))219        if created:220            st.success(f"New features: {', '.join(created)}")221        else:222            st.info("No MODIS band columns detected. Statistical aggregates added.")223        st.dataframe(df_train_eng[num_features_eng].head(8), use_container_width=True)224 225with tab3:226    if not st.session_state["data_loaded"]:227        st.info("Load the dataset first.")228    elif "num_features_eng" not in st.session_state:229        st.info("Visit the Features tab first.")230    elif not train_btn and not st.session_state["model_trained"]:231        st.info("Click Train Model in the sidebar.")232    else:233        import numpy as np234        import matplotlib.pyplot as plt235        import seaborn as sns236        import pandas as pd237        from sklearn.model_selection import train_test_split238        from sklearn.preprocessing import LabelEncoder239        from sklearn.ensemble import RandomForestClassifier240        from sklearn.metrics import accuracy_score, f1_score, classification_report, confusion_matrix241        df_train_eng = st.session_state["df_train_eng"]242        df_test_eng = st.session_state["df_test_eng"]243        TARGET = st.session_state["TARGET"]244        num_features_eng = st.session_state["num_features_eng"]245        if train_btn:246            with st.spinner(f"Training {selected_model}..."):247                le = LabelEncoder()248                le.fit(df_train_eng[TARGET])249                y_train_full = le.transform(df_train_eng[TARGET])250                known_set = set(le.classes_)251                y_test_all = np.array([le.transform([c])[0] if c in known_set else -1 for c in df_test_eng[TARGET].values])252                X_train_full = df_train_eng[num_features_eng].fillna(df_train_eng[num_features_eng].median())253                X_test_all = df_test_eng[num_features_eng].fillna(df_test_eng[num_features_eng].median())254                mask = y_test_all != -1255                X_test_ev = X_test_all[mask]256                y_test_ev = y_test_all[mask]257                X_tr, X_val, y_tr, y_val = train_test_split(X_train_full, y_train_full, test_size=test_size, random_state=42, stratify=y_train_full)258                if selected_model == "Random Forest":259                    model = RandomForestClassifier(n_estimators=200, max_depth=12, random_state=42, n_jobs=-1)260                elif selected_model == "XGBoost":261                    import xgboost as xgb262                    model = xgb.XGBClassifier(n_estimators=200, max_depth=6, learning_rate=0.1, eval_metric="mlogloss", random_state=42, n_jobs=-1, verbosity=0)263                else:264                    import lightgbm as lgb265                    model = lgb.LGBMClassifier(n_estimators=300, max_depth=8, learning_rate=0.05, num_leaves=63, random_state=42, n_jobs=-1, verbose=-1)266                model.fit(X_tr, y_tr)267                val_acc = accuracy_score(y_val, model.predict(X_val))268                val_f1 = f1_score(y_val, model.predict(X_val), average="weighted", zero_division=0)269                model.fit(X_train_full, y_train_full)270                y_pred = model.predict(X_test_ev)271                test_acc = accuracy_score(y_test_ev, y_pred)272                test_f1 = f1_score(y_test_ev, y_pred, average="weighted", zero_division=0)273            st.session_state.update({274                "model": model, "X_test_ev": X_test_ev, "y_test_ev": y_test_ev,275                "y_pred": y_pred, "le": le, "val_acc": val_acc, "val_f1": val_f1,276                "test_acc": test_acc, "test_f1": test_f1,277                "model_trained": True, "selected_model": selected_model,278            })279        if st.session_state["model_trained"]:280            PALETTE = ["#64ffda","#f07178","#c3e88d","#82aaff","#ffcb6b","#ff5572"]281            st.success(f"{st.session_state['selected_model']} trained!")282            c1, c2, c3, c4 = st.columns(4)283            c1.metric("Val Accuracy", f"{st.session_state['val_acc']:.4f}")284            c2.metric("Val F1", f"{st.session_state['val_f1']:.4f}")285            c3.metric("Test Accuracy", f"{st.session_state['test_acc']:.4f}")286            c4.metric("Test F1", f"{st.session_state['test_f1']:.4f}")287            st.markdown("---")288            le = st.session_state["le"]289            y_test_ev = st.session_state["y_test_ev"]290            y_pred = st.session_state["y_pred"]291            st.subheader("Classification Report")292            report = classification_report(y_test_ev, y_pred, target_names=[str(c) for c in le.inverse_transform(np.unique(y_test_ev))], zero_division=0, output_dict=True)293            st.dataframe(pd.DataFrame(report).T.round(3), use_container_width=True)294            st.markdown("---")295            st.subheader("Confusion Matrix")296            ul = np.unique(np.concatenate([y_test_ev, y_pred]))297            lnames = [str(c) for c in le.inverse_transform(ul)]298            cm = confusion_matrix(y_test_ev, y_pred, labels=ul)299            cm_pct = cm.astype(float) / (cm.sum(axis=1, keepdims=True) + 1e-9) * 100300            fig, ax = plt.subplots(figsize=(max(6, len(ul) * 0.65), max(5, len(ul) * 0.55)))301            sns.heatmap(cm_pct, annot=True, fmt=".1f", cmap="RdYlGn", vmin=0, vmax=100, xticklabels=lnames, yticklabels=lnames, linewidths=0.4, ax=ax)302            ax.set_title(f"{st.session_state['selected_model']} - Confusion Matrix")303            ax.set_xlabel("Predicted")304            ax.set_ylabel("True")305            plt.tight_layout()306            st.pyplot(fig)307            plt.close()308 309with tab4:310    if not HAS_SHAP:311        st.warning("SHAP is not installed.")312    elif not st.session_state["model_trained"]:313        st.info("Train a model first in the Model tab.")314    elif not run_xai:315        st.warning("Enable Run SHAP in the sidebar, then retrain.")316    else:317        import numpy as np318        import matplotlib.pyplot as plt319        import pandas as pd320        import shap321        PALETTE = ["#64ffda","#f07178","#c3e88d","#82aaff","#ffcb6b","#ff5572"]322        model_xai = st.session_state["model"]323        X_ev = st.session_state["X_test_ev"]324        feat_names = st.session_state["num_features_eng"]325        mname = st.session_state["selected_model"]326        shap_n = st.slider("Samples for SHAP", 50, 200, 100, 50)327        with st.spinner("Computing SHAP values..."):328            explainer = shap.TreeExplainer(model_xai)329            samp = X_ev.iloc[:shap_n] if hasattr(X_ev, "iloc") else X_ev[:shap_n]330            shap_vals = explainer.shap_values(samp)331            sv_arr = np.array(shap_vals)332            if sv_arr.ndim == 2:333                sm = np.abs(sv_arr).mean(axis=0)334            elif sv_arr.ndim == 3:335                sm = np.abs(sv_arr).mean(axis=1).mean(axis=0)336            else:337                sm = np.abs(sv_arr).reshape(-1, len(feat_names)).mean(axis=0)338            if sm.shape[0] != len(feat_names):339                sm = np.abs(sv_arr).reshape(len(feat_names), -1).mean(axis=1)340        st.success("SHAP computed!")341        imp = pd.Series(sm, index=feat_names).sort_values(ascending=False)342        top = min(20, len(imp))343        fig, ax = plt.subplots(figsize=(10, top * 0.42))344        ax.barh(imp.index[:top][::-1], imp.values[:top][::-1], color=PALETTE[0])345        ax.set_title("SHAP - Mean |SHAP value|")346        ax.set_xlabel("Mean |SHAP|")347        ax.grid(axis="x", alpha=0.3)348        plt.tight_layout()349        st.pyplot(fig)350        plt.close()351        if hasattr(model_xai, "feature_importances_"):352            imp_bi = pd.Series(model_xai.feature_importances_, index=feat_names).sort_values(ascending=False)353            top_n = min(20, len(imp_bi))354            fig, ax = plt.subplots(figsize=(10, top_n * 0.42))355            ax.barh(imp_bi.index[:top_n][::-1], imp_bi.values[:top_n][::-1], color=PALETTE[2])356            ax.set_title(f"{mname} - Feature Importances")357            ax.set_xlabel("Importance")358            ax.grid(axis="x", alpha=0.3)359            plt.tight_layout()360            st.pyplot(fig)361            plt.close()