lawn-estimator-dev / scripts /exp4 /train_eomt.py
TempuraML's picture
perf(exp4): trainer memory + visibility — freeze layers, batch 1, step logging, mid-epoch ckpt
a46e830
Raw
History Blame Contribute Delete
14.9 kB
"""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()