Spaces:
Sleeping
Sleeping
| """Spectral baseline classifier: KNN on raw 7-band pixel values, no spatial context.""" | |
| import numpy as np | |
| from config import NUM_CHANNELS, NUM_CLASSES, IGNORE_INDEX | |
| from metrics import compute_metrics, metrics_markdown | |
| def _knn_predict( | |
| train_X: np.ndarray, | |
| train_y: np.ndarray, | |
| query_X: np.ndarray, | |
| k: int, | |
| chunk: int = 50_000, | |
| ) -> np.ndarray: | |
| """Chunked nearest-neighbour prediction to keep peak RAM reasonable.""" | |
| N = len(query_X) | |
| preds = np.empty(N, dtype=np.int64) | |
| k = min(k, len(train_X)) | |
| for start in range(0, N, chunk): | |
| end = min(start + chunk, N) | |
| block = query_X[start:end] # (B, 7) | |
| dists = np.sum((block[:, None, :] - train_X[None, :, :]) ** 2, axis=2) # (B, N_tr) | |
| nn_idx = np.argpartition(dists, k - 1, axis=1)[:, :k] # (B, k) | |
| labels = train_y[nn_idx] # (B, k) | |
| if k == 1: | |
| preds[start:end] = labels[:, 0] | |
| else: | |
| # Vectorised majority vote | |
| votes = (labels[:, :, None] == np.arange(NUM_CLASSES)[None, None, :]).sum(axis=1) | |
| preds[start:end] = votes.argmax(axis=1) | |
| return preds | |
| def run_knn_baseline( | |
| full_image: np.ndarray, | |
| full_train_mask: np.ndarray, | |
| full_val_mask: np.ndarray, | |
| val_images: np.ndarray, | |
| k: int = 3, | |
| ): | |
| """ | |
| Train KNN on labeled training pixels; predict (a) the full scene and (b) each | |
| validation patch. Evaluate against the full validation mask. | |
| Returns | |
| ------- | |
| full_pred : (H, W) – class index for every pixel in the scene | |
| val_preds : (N, ph, pw) – patch-level predictions for step-4 comparison | |
| metrics : dict | |
| metrics_md : str | |
| """ | |
| C, H, W = full_image.shape | |
| labeled = full_train_mask != IGNORE_INDEX | |
| if not labeled.any(): | |
| raise ValueError("No labeled training pixels found in TRAINING.tif.") | |
| train_X = full_image[:, labeled].T # (N_tr, 7) | |
| train_y = full_train_mask[labeled] # (N_tr,) | |
| # --- Full scene prediction --- | |
| all_X = full_image.reshape(C, H * W).T # (H*W, 7) | |
| full_pred = _knn_predict(train_X, train_y, all_X, k).reshape(H, W) | |
| # --- Validation patch predictions (same patches as UNet) --- | |
| N_val, ph, pw = val_images.shape[0], val_images.shape[2], val_images.shape[3] | |
| all_patch_X = np.concatenate( | |
| [p.reshape(C, -1).T for p in val_images], axis=0 | |
| ) # (N_val * ph * pw, 7) | |
| patch_preds_flat = _knn_predict(train_X, train_y, all_patch_X, k) | |
| val_preds = patch_preds_flat.reshape(N_val, ph, pw).astype(np.int64) | |
| # --- Metrics on full val mask --- | |
| metrics = compute_metrics(full_pred.ravel(), full_val_mask.ravel()) | |
| metrics_md = metrics_markdown(metrics, title=f"KNN Baseline (k={k})") | |
| return full_pred.astype(np.int64), val_preds, metrics, metrics_md | |