segspace_app / baseline.py
functionNormally
Redesign: five-step pedagogical flow with spectral baseline
089078d
Raw
History Blame Contribute Delete
2.93 kB
"""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