osky9/Land_Cover_classification
0
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()