#!/usr/bin/env bash # SUM Parts - let inference rule out classes we know are absent # # The network is a closed-set classifier: 13 logits, softmax, argmax. There is # no "none of these". A point whose true class was never in the training # vocabulary still gets a label - whichever learned concept sits closest in # feature space. # # On the Seosan road tiles that produces 21% water and 5% boat. Neither exists # there. What the model learned as "water" in Helsinki - dark, flat, smooth, # horizontal - describes asphalt exactly, and boats are what sits on water, so # the hallucination is internally consistent. # # We know those classes are absent. The model does not. Masking their logits # before the argmax hands each of those points to its runner-up class instead. # # This is not a fix for the domain gap; it is telling the model something we # already know. Whether it helps depends entirely on what the runner-up is, # which is why it is worth measuring rather than assuming. # # Usage: # bash scripts/patch_class_mask.sh # SUMPARTS_MASK_CLASSES=4,6 # 4=water, 6=boat # # Idempotent. set -euo pipefail source "$HOME/miniconda3/etc/profile.d/conda.sh" conda activate sumparts REPO="${1:-$HOME/sum-parts/semantic_segmentation/PointNeXt_bundle}" MAIN="$REPO/examples/segmentation/main.py" [ -f "$MAIN" ] || { echo "error: $MAIN not found" >&2; exit 1; } if grep -q 'SUMPARTS-CLASS-MASK' "$MAIN"; then echo "already patched" exit 0 fi cp -n "$MAIN" "$MAIN.orig" 2>/dev/null || true python - "$MAIN" <<'PY' import sys from pathlib import Path p = Path(sys.argv[1]) src = p.read_text(encoding="utf-8") old = """ pred = all_logits.argmax(dim=1) if label is not None: cm.update(pred, label)""" new = """ # SUMPARTS-CLASS-MASK: drop classes we know cannot occur in this scene # before taking the argmax, so their points fall through to the # runner-up instead. Set SUMPARTS_MASK_CLASSES to a comma-separated # list of class indices, e.g. "4,6" for water and boat. _mask = os.environ.get('SUMPARTS_MASK_CLASSES', '').strip() if _mask: _idx = [int(x) for x in _mask.split(',') if x.strip() != ''] if cloud_idx == 0: logging.info(f' masking classes {_idx} out of the argmax') all_logits[:, _idx] = float('-inf') pred = all_logits.argmax(dim=1) if label is not None: cm.update(pred, label)""" if old not in src: print("PATTERN NOT FOUND -- main.py differs from what this patch expects", file=sys.stderr) raise SystemExit(1) p.write_text(src.replace(old, new), encoding="utf-8") print("patched:", p) PY python -c "import ast,sys; ast.parse(open(sys.argv[1], encoding='utf-8').read())" "$MAIN" \ && echo "syntax OK" echo "PATCH DONE (original kept at $MAIN.orig)"