Spaces:
Running
Running
| """Exp 4 Phase 3 — EoMT fine-tune on auto-labeled DOGIS crops. RUNS ON THE M5 MAC. | |
| Fine-tunes tue-mps EoMT (DINOv3 backbone; --base-model swaps in the DINOv2/ | |
| Apache fallback) on the crop dataset produced by export_training_crops.py + | |
| autolabel_sam3.py. Our 10-class taxonomy replaces the ADE head | |
| (`ignore_mismatched_sizes`); dormant and shadowed turf are first-class LAWN. | |
| Overnight-safe: checkpoints every epoch (latest + best-mIoU), resumes from | |
| --out automatically (model + optimizer + epoch), never prompts, prints one | |
| parseable line per epoch and a REPORT BEGIN/END block at the end (also on | |
| crash). Copy the whole block back. | |
| Setup on a fresh Mac (same venv as the labeler): | |
| python3 -m venv ~/lawn-train && source ~/lawn-train/bin/activate | |
| pip install torch transformers pillow numpy | |
| python train_eomt.py --data <crops folder> --out ~/lawn-train/run1 | |
| Split is BY CLUSTER (LiDAR-tile neighborhoods) so val parcels are spatially | |
| disjoint from train. Crops whose id is listed in --exclude-file (one id per | |
| line — the owner's "fix" verdicts until they're corrected) are dropped. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import hashlib | |
| import json | |
| import platform | |
| import random | |
| import time | |
| import traceback | |
| from pathlib import Path | |
| import numpy as np | |
| from PIL import Image, ImageEnhance | |
| BASE_MODEL = "tue-mps/eomt-dinov3-ade-semantic-large-512" | |
| BASE_REVISION = "1d1f172700ab69371df7afe778fe133fff7dd521" # = segmentation.MODEL_REVISIONS | |
| CLASSES = ["ignore", "lawn", "tree", "building", "driveway", "sidewalk", | |
| "dirt", "pool", "road", "vehicle"] | |
| IGNORE_INDEX = 0 # our class 0; remapped to 255 for the loss | |
| INPUT = 512 | |
| # ------------------------------------------------------------------- dataset | |
| class CropDataset: | |
| def __init__(self, data: Path, ids: list[str], train: bool, seed: int = 0): | |
| self.data, self.ids, self.train = data, ids, train | |
| self.rng = random.Random(seed) | |
| def __len__(self): | |
| return len(self.ids) | |
| def item(self, i: int) -> tuple[np.ndarray, np.ndarray]: | |
| cid = self.ids[i] | |
| img = Image.open(self.data / "images" / f"{cid}.jpg").convert("RGB") | |
| lab = Image.open(self.data / "labels" / f"{cid}.png") | |
| if self.train: | |
| img, lab = self.augment(img, lab) | |
| img = img.resize((INPUT, INPUT), Image.BILINEAR) | |
| lab = lab.resize((INPUT, INPUT), Image.NEAREST) | |
| sem = np.asarray(lab, dtype=np.int64) # values 0..N-1; 0 = ignore | |
| return np.asarray(img), sem | |
| def augment(self, img: Image.Image, lab: Image.Image): | |
| r = self.rng | |
| # random scale-crop: the model must work from ~60 m context (Google tile | |
| # scale at inference) up to near-native detail | |
| w, h = img.size | |
| s = r.uniform(0.5, 1.0) | |
| cw = int(w * s) | |
| x0, y0 = r.randint(0, w - cw), r.randint(0, h - cw) | |
| img, lab = img.crop((x0, y0, x0 + cw, y0 + cw)), lab.crop((x0, y0, x0 + cw, y0 + cw)) | |
| if r.random() < 0.5: | |
| img, lab = img.transpose(Image.FLIP_LEFT_RIGHT), lab.transpose(Image.FLIP_LEFT_RIGHT) | |
| k = r.choice([0, 1, 2, 3]) | |
| for _ in range(k): | |
| img, lab = img.transpose(Image.ROTATE_90), lab.transpose(Image.ROTATE_90) | |
| # photometric bridge between leaf-off training and leaf-on inference: | |
| # wide color/saturation jitter BOTH directions (dormant<->green) | |
| img = ImageEnhance.Color(img).enhance(r.uniform(0.5, 1.7)) | |
| img = ImageEnhance.Brightness(img).enhance(r.uniform(0.75, 1.25)) | |
| img = ImageEnhance.Contrast(img).enhance(r.uniform(0.8, 1.25)) | |
| return img, lab | |
| def set_trainable(model, freeze_layers: int) -> tuple[float, float]: | |
| """Freeze embeddings + the first `freeze_layers` transformer blocks; train the | |
| rest and all heads. EoMT is encoder-only (the 24 shared layers hold ~96% of the | |
| params), so freezing the lower blocks is the main memory/speed lever. Returns | |
| (trainable_M, total_M).""" | |
| import re | |
| for name, p in model.named_parameters(): | |
| frozen = name.startswith(("embeddings", "rope_embeddings")) | |
| m = re.search(r"layers\.(\d+)\.", name) | |
| if m and int(m.group(1)) < freeze_layers: | |
| frozen = True | |
| p.requires_grad = not frozen | |
| tr = sum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6 | |
| tot = sum(p.numel() for p in model.parameters()) / 1e6 | |
| return tr, tot | |
| def split_ids(data: Path, exclude: set[str]) -> tuple[list[str], list[str]]: | |
| """90/10 by cluster hash — spatially disjoint val.""" | |
| train, val = [], [] | |
| for p in sorted((data / "labels").glob("*.png")): | |
| cid = p.stem | |
| if cid in exclude: | |
| continue | |
| cluster = json.loads((data / "meta" / f"{cid}.json").read_text())["cluster"] | |
| h = int(hashlib.sha1(str(cluster).encode()).hexdigest(), 16) % 10 | |
| (val if h == 0 else train).append(cid) | |
| return train, val | |
| # --------------------------------------------------------------------- train | |
| def main() -> None: | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--data", required=True) | |
| ap.add_argument("--out", required=True) | |
| ap.add_argument("--base-model", default=BASE_MODEL) | |
| ap.add_argument("--revision", default=BASE_REVISION) | |
| ap.add_argument("--epochs", type=int, default=20) | |
| ap.add_argument("--batch", type=int, default=1, | |
| help="1 keeps memory low on 16 GB Macs; raise only if RAM allows") | |
| ap.add_argument("--lr", type=float, default=2e-5) | |
| ap.add_argument("--exclude-file", default=None) | |
| ap.add_argument("--seed", type=int, default=0) | |
| ap.add_argument("--freeze-layers", type=int, default=18, | |
| help="freeze the first N of 24 transformer blocks (+embeddings) to cut " | |
| "optimizer/gradient memory; 0 = full fine-tune, 24 = heads only") | |
| ap.add_argument("--log-every", type=int, default=20, help="print step progress every N steps") | |
| ap.add_argument("--ckpt-every", type=int, default=400, help="mid-epoch checkpoint every N steps") | |
| ap.add_argument("--max-steps", type=int, default=0, help="smoke: stop after N steps") | |
| args = ap.parse_args() | |
| out = Path(args.out) | |
| out.mkdir(parents=True, exist_ok=True) | |
| data = Path(args.data) | |
| print("=" * 66) | |
| print("TRAIN REPORT BEGIN — copy everything down to REPORT END") | |
| print("=" * 66) | |
| print(f"machine: {platform.machine()} | {platform.platform()}") | |
| try: | |
| import torch | |
| from transformers import AutoImageProcessor, AutoModelForUniversalSegmentation | |
| device = ("mps" if torch.backends.mps.is_available() | |
| else "cuda" if torch.cuda.is_available() else "cpu") | |
| torch.manual_seed(args.seed) | |
| print(f"torch {torch.__version__} | device: {device} | base: {args.base_model}") | |
| exclude = set() | |
| if args.exclude_file: | |
| exclude = {line.strip() for line in Path(args.exclude_file).read_text().splitlines() | |
| if line.strip()} | |
| train_ids, val_ids = split_ids(data, exclude) | |
| print(f"dataset: {len(train_ids)} train / {len(val_ids)} val " | |
| f"({len(exclude)} excluded)") | |
| if not train_ids or not val_ids: | |
| raise SystemExit("empty split — need labeled crops from >=2 clusters") | |
| processor = AutoImageProcessor.from_pretrained(args.base_model, revision=args.revision) | |
| id2label = dict(enumerate(CLASSES)) | |
| model = AutoModelForUniversalSegmentation.from_pretrained( | |
| args.base_model, revision=args.revision, | |
| id2label=id2label, label2id={v: k for k, v in id2label.items()}, | |
| ignore_mismatched_sizes=True).to(device) | |
| tr_m, tot_m = set_trainable(model, args.freeze_layers) | |
| print(f"trainable: {tr_m:.0f}M / {tot_m:.0f}M params " | |
| f"({args.freeze_layers} of 24 blocks frozen)") | |
| optim = torch.optim.AdamW( | |
| [p for p in model.parameters() if p.requires_grad], lr=args.lr, weight_decay=1e-4) | |
| start_epoch, best_miou = 0, -1.0 | |
| state_path = out / "latest" / "train_state.json" | |
| if state_path.exists(): | |
| st = json.loads(state_path.read_text()) | |
| start_epoch, best_miou = st["epoch"] + 1, st["best_miou"] | |
| model = AutoModelForUniversalSegmentation.from_pretrained(out / "latest").to(device) | |
| set_trainable(model, args.freeze_layers) | |
| opt_file = out / "latest" / "optimizer.pt" | |
| optim = torch.optim.AdamW( | |
| [p for p in model.parameters() if p.requires_grad], lr=args.lr, weight_decay=1e-4) | |
| if opt_file.exists(): | |
| optim.load_state_dict(torch.load(opt_file, map_location=device)) | |
| print(f"RESUMED from epoch {start_epoch} (best mIoU {best_miou:.3f})") | |
| sched = torch.optim.lr_scheduler.CosineAnnealingLR( | |
| optim, T_max=max(args.epochs, 1), last_epoch=start_epoch - 1) | |
| train_ds = CropDataset(data, train_ids, train=True, seed=args.seed) | |
| val_ds = CropDataset(data, val_ids, train=False) | |
| def batches(ds: CropDataset, batch: int, shuffle: bool): | |
| # Build mask_labels/class_labels explicitly (skip ignore=0) instead of | |
| # processor(segmentation_maps=...): EomtImageProcessor mangles the ignore | |
| # value (255 -> a class index -> "index 255 out of bounds" in the loss). | |
| order = list(range(len(ds))) | |
| if shuffle: | |
| random.Random(args.seed + start_epoch).shuffle(order) | |
| for i in range(0, len(order), batch): | |
| items = [ds.item(j) for j in order[i:i + batch]] | |
| imgs = [it[0] for it in items] | |
| enc = processor(images=imgs, return_tensors="pt", do_resize=False) | |
| mask_labels, class_labels = [], [] | |
| for _, sem in items: | |
| present = [int(c) for c in np.unique(sem) if int(c) != IGNORE_INDEX] | |
| masks = (np.stack([sem == c for c in present]).astype("float32") | |
| if present else np.zeros((0, *sem.shape), "float32")) | |
| mask_labels.append(torch.from_numpy(masks)) | |
| class_labels.append(torch.tensor(present, dtype=torch.int64)) | |
| yield {"pixel_values": enc["pixel_values"], | |
| "mask_labels": mask_labels, "class_labels": class_labels} | |
| def to_dev(enc): | |
| return {k: (v.to(device) if hasattr(v, "to") else | |
| [t.to(device) for t in v] if isinstance(v, list) else v) | |
| for k, v in enc.items()} | |
| steps_per_epoch = (len(train_ds) + args.batch - 1) // args.batch | |
| print(f"{steps_per_epoch} steps/epoch (batch {args.batch}) — progress every " | |
| f"{args.log_every}, checkpoint every {args.ckpt_every}", flush=True) | |
| def save_latest(ep): | |
| model.save_pretrained(out / "latest") | |
| processor.save_pretrained(out / "latest") | |
| torch.save(optim.state_dict(), out / "latest" / "optimizer.pt") | |
| state_path.write_text(json.dumps({"epoch": ep, "best_miou": best_miou})) | |
| step_count = 0 | |
| for epoch in range(start_epoch, args.epochs): | |
| model.train() | |
| t0, losses, se = time.time(), [], 0 | |
| for enc in batches(train_ds, args.batch, shuffle=True): | |
| enc = to_dev(enc) | |
| outp = model(**enc) | |
| loss = outp.loss | |
| loss.backward() | |
| optim.step() | |
| optim.zero_grad(set_to_none=True) | |
| losses.append(float(loss.detach().cpu())) | |
| step_count += 1 | |
| se += 1 | |
| if args.log_every and se % args.log_every == 0: | |
| sps = (time.time() - t0) / se | |
| print(f" e{epoch} step {se}/{steps_per_epoch} | " | |
| f"loss {np.mean(losses[-args.log_every:]):.3f} | " | |
| f"{sps:.1f}s/step | epoch ETA {sps * (steps_per_epoch - se) / 60:.0f}m", | |
| flush=True) | |
| if args.ckpt_every and se % args.ckpt_every == 0: | |
| save_latest(epoch - 1) # mid-epoch: keep prior epoch as the resume point | |
| print(f" e{epoch} checkpoint at step {se}", flush=True) | |
| if args.max_steps and step_count >= args.max_steps: | |
| print(f"SMOKE STOP after {step_count} steps | " | |
| f"losses: {[round(x, 3) for x in losses]}") | |
| print("REPORT END") | |
| return | |
| sched.step() | |
| # validation mIoU (semantic argmax vs label maps) | |
| model.eval() | |
| inter = np.zeros(len(CLASSES)) | |
| union = np.zeros(len(CLASSES)) | |
| with torch.inference_mode(): | |
| for i in range(len(val_ds)): | |
| img_arr, sem_true = val_ds.item(i) | |
| enc = to_dev(processor(images=[img_arr], return_tensors="pt", do_resize=False)) | |
| outp = model(**enc) | |
| pred = processor.post_process_semantic_segmentation( | |
| outp, target_sizes=[sem_true.shape])[0] | |
| pred = pred.cpu().numpy() | |
| valid = sem_true != IGNORE_INDEX | |
| for c in range(1, len(CLASSES)): | |
| pi, ti = pred == c, sem_true == c | |
| inter[c] += (pi & ti & valid).sum() | |
| union[c] += ((pi | ti) & valid).sum() | |
| ious = {CLASSES[c]: round(inter[c] / union[c], 3) | |
| for c in range(1, len(CLASSES)) if union[c] > 0} | |
| miou = float(np.mean(list(ious.values()))) if ious else 0.0 | |
| model.save_pretrained(out / "latest") | |
| processor.save_pretrained(out / "latest") | |
| torch.save(optim.state_dict(), out / "latest" / "optimizer.pt") | |
| state_path.write_text(json.dumps({"epoch": epoch, "best_miou": max(best_miou, miou)})) | |
| if miou > best_miou: | |
| best_miou = miou | |
| model.save_pretrained(out / "best") | |
| processor.save_pretrained(out / "best") | |
| print(f"EPOCH {epoch}: loss {np.mean(losses):.3f} | mIoU {miou:.3f} " | |
| f"| per-class {json.dumps(ious)} | {time.time() - t0:.0f}s " | |
| f"| best {best_miou:.3f}{' *NEW BEST*' if miou == best_miou else ''}") | |
| print(f"\nDONE: best mIoU {best_miou:.3f} -> {out / 'best'}") | |
| print("Next: run the Mac-side eval + copy `best/` back for the showdown.") | |
| except SystemExit as e: | |
| print(e) | |
| except Exception: | |
| print("TRAINING FAILED:") | |
| print(traceback.format_exc()) | |
| print("=" * 66) | |
| print("REPORT END") | |
| print("=" * 66) | |
| if __name__ == "__main__": | |
| main() | |