"""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