Instructions to use m1sc/reach-down-vit-release with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- DepthAnythingV2
How to use m1sc/reach-down-vit-release with DepthAnythingV2:
# Install from https://github.com/DepthAnything/Depth-Anything-V2 # Load the model and infer depth from an image import cv2 import torch from huggingface_hub import hf_hub_download from depth_anything_v2.dpt import DepthAnythingV2 # instantiate the model model = DepthAnythingV2(encoder="<ENCODER>", features=<NUMBER_OF_FEATURES>, out_channels=<OUT_CHANNELS>) # load the weights filepath = hf_hub_download(repo_id="m1sc/reach-down-vit-release", filename="depth_anything_v2_<ENCODER>.pth", repo_type="model") state_dict = torch.load(filepath, map_location="cpu") model.load_state_dict(state_dict) model.eval() raw_img = cv2.imread("your/image/path") depth = model.infer_image(raw_img) # HxW raw depth map in numpy - Notebooks
- Google Colab
- Kaggle
Download scripts/train.py from m1sc/reach-down-vit-release: direct link, hf CLI and curl.
- Browser
- Download file 7.67 kB
-
https://huggingface.co/m1sc/reach-down-vit-release/resolve/main/scripts/train.py
- Command line
-
hf download hf://m1sc/reach-down-vit-release/scripts/train.py
-
curl -L -o train.py https://huggingface.co/m1sc/reach-down-vit-release/resolve/main/scripts/train.py
7.67 kB
| #!/usr/bin/env python3 | |
| """Fine-tune the landing-value ViT on HJ-labeled tiles. | |
| .venv/bin/python scripts/train.py --data data/tiles --epochs 8 | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader, Subset | |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) | |
| from reachdown.paths import data_path | |
| from reachdown.data import TileDataset | |
| from reachdown.model import LandingValueViT | |
| def soft_dice_loss(logits: torch.Tensor, target: torch.Tensor, eps: float = 1.0) -> torch.Tensor: | |
| p = torch.sigmoid(logits) | |
| num = 2.0 * (p * target).sum(dim=(1, 2, 3)) + eps | |
| den = (p + target).sum(dim=(1, 2, 3)) + eps | |
| return 1.0 - (num / den).mean() | |
| def auroc(scores: np.ndarray, labels: np.ndarray) -> float: | |
| """Rank-based AUROC (labels thresholded at 0.5); subsampled for speed.""" | |
| idx = np.random.default_rng(0).choice(scores.size, min(scores.size, 200_000), replace=False) | |
| s, y = scores.ravel()[idx], labels.ravel()[idx] > 0.5 | |
| n_pos, n_neg = int(y.sum()), int((~y).sum()) | |
| if not n_pos or not n_neg: | |
| return float("nan") | |
| ranks = s.argsort().argsort().astype(np.float64) + 1 | |
| return float((ranks[y].sum() - n_pos * (n_pos + 1) / 2) / (n_pos * n_neg)) | |
| def evaluate(model, loader, device) -> tuple[float, float]: | |
| model.eval() | |
| scores, labels, maes = [], [], [] | |
| for x, y in loader: | |
| x, y = x.to(device), y.to(device) | |
| p = torch.sigmoid(model(x)) | |
| scores.append(p.cpu().numpy()) | |
| labels.append(y.cpu().numpy()) | |
| maes.append(float((p - y).abs().mean())) | |
| return auroc(np.concatenate(scores), np.concatenate(labels)), float(np.mean(maes)) | |
| def main() -> None: | |
| ap = argparse.ArgumentParser(description=__doc__) | |
| ap.add_argument("--data", type=Path, default=data_path("data", "tiles")) | |
| ap.add_argument("--label-key", default="label_value", choices=("label_value", "label_run")) | |
| ap.add_argument("--safety-buffer", type=float, default=0.0, | |
| help="re-compose conservative labels from stored margin_hj (m)") | |
| ap.add_argument("--input-size", type=int, default=364, help="multiple of 14") | |
| ap.add_argument("--epochs", type=int, default=8) | |
| ap.add_argument("--batch-size", type=int, default=8) | |
| ap.add_argument("--lr", type=float, default=3e-4, help="decoder LR; backbone gets 0.1x") | |
| ap.add_argument("--false-safe-weight", type=float, default=1.0, | |
| help=">1 penalizes predicting-safe-where-unsafe (conservative bias)") | |
| ap.add_argument("--val-frac", type=float, default=0.2) | |
| ap.add_argument("--seed", type=int, default=0, help="torch/numpy seed for multi-seed runs") | |
| ap.add_argument("--from-scratch", action="store_true", | |
| help="FW-3: random-init backbone (no depth pretraining)") | |
| ap.add_argument("--unfreeze-blocks", type=int, default=4, | |
| help="FW-4: number of last encoder blocks to unfreeze") | |
| ap.add_argument("--zero-channels", type=int, nargs="*", default=[], | |
| help="FW-1: input channels to zero (e.g. 4 = analytic margin)") | |
| ap.add_argument("--out", type=Path, default=Path("runs/vit")) | |
| args = ap.parse_args() | |
| torch.manual_seed(args.seed) | |
| np.random.seed(args.seed) | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| dataset = TileDataset(args.data, input_size=args.input_size, label_key=args.label_key, | |
| safety_buffer_m=args.safety_buffer, | |
| zero_channels=tuple(args.zero_channels)) | |
| # split by terrain, not tile: tiles differing only in aircraft state must | |
| # not straddle the split | |
| terrains = [dataset.terrain_seed(i) for i in range(len(dataset))] | |
| unique = sorted(set(terrains)) | |
| if len(unique) < 2: | |
| raise SystemExit("need >= 2 distinct terrains for a terrain-split validation") | |
| val_terrains = set(unique[max(1, int(len(unique) * (1 - args.val_frac))):]) | |
| val_idx = [i for i, t in enumerate(terrains) if t in val_terrains] | |
| train_idx = [i for i, t in enumerate(terrains) if t not in val_terrains] | |
| train_set, val_set = Subset(dataset, train_idx), Subset(dataset, val_idx) | |
| train_loader = DataLoader( | |
| train_set, batch_size=args.batch_size, shuffle=True, num_workers=4, | |
| persistent_workers=True, pin_memory=True, | |
| ) | |
| val_loader = DataLoader( | |
| val_set, batch_size=args.batch_size, num_workers=2, | |
| persistent_workers=True, pin_memory=True, | |
| ) | |
| model = LandingValueViT(unfreeze_last_blocks=args.unfreeze_blocks, | |
| pretrained=not args.from_scratch).to(device) | |
| backbone_params, decoder_params = model.trainable_parameters() | |
| opt = torch.optim.AdamW( | |
| [{"params": decoder_params, "lr": args.lr}, | |
| {"params": backbone_params, "lr": 0.1 * args.lr}], | |
| weight_decay=1e-4, | |
| ) | |
| sched = torch.optim.lr_scheduler.CosineAnnealingLR( | |
| opt, T_max=args.epochs * len(train_loader) | |
| ) | |
| args.out.mkdir(parents=True, exist_ok=True) | |
| # persist the exact run configuration: reproducibility surface for the paper | |
| import json | |
| (args.out / "args.json").write_text(json.dumps( | |
| {k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items()}, | |
| indent=2, sort_keys=True)) | |
| history = args.out / "history.csv" | |
| history.write_text("epoch,train_loss,val_auroc,val_mae\n") | |
| best = -float("inf") | |
| print(f"{len(train_set)} train / {len(val_set)} val tiles on {device}; " | |
| f"trainable params: {sum(p.numel() for p in decoder_params + backbone_params) / 1e6:.1f}M") | |
| for epoch in range(args.epochs): | |
| model.train() | |
| t0, losses = time.time(), [] | |
| for x, y in train_loader: | |
| x, y = x.to(device), y.to(device) | |
| with torch.autocast(device, dtype=torch.bfloat16, enabled=device == "cuda"): | |
| logits = model(x) | |
| # asymmetric BCE: a false-safe (predict safe where label unsafe) | |
| # costs more than a false-unsafe, biasing the map conservative | |
| w = 1.0 + (args.false_safe_weight - 1.0) * (y < 0.5).float() | |
| bce = F.binary_cross_entropy_with_logits(logits, y, weight=w) | |
| loss = bce + 0.5 * soft_dice_loss(logits, y) | |
| opt.zero_grad(set_to_none=True) | |
| loss.backward() | |
| opt.step() | |
| sched.step() | |
| losses.append(loss.item()) | |
| val_auroc, val_mae = evaluate(model, val_loader, device) | |
| print(f"epoch {epoch + 1:02d}/{args.epochs} loss {np.mean(losses):.4f} " | |
| f"val AUROC {val_auroc:.4f} val MAE {val_mae:.4f} ({time.time() - t0:.0f}s)") | |
| with history.open("a") as f: | |
| f.write(f"{epoch + 1},{np.mean(losses):.6f},{val_auroc:.6f},{val_mae:.6f}\n") | |
| # single-class validation yields NaN AUROC; fall back to MAE so a | |
| # checkpoint is always produced | |
| score = -val_mae if np.isnan(val_auroc) else val_auroc | |
| if score > best: | |
| best = score | |
| torch.save( | |
| {"state_dict": model.state_dict(), "input_size": args.input_size, | |
| "label_key": args.label_key, "val_auroc": val_auroc, | |
| "zero_channels": list(args.zero_channels), | |
| "unfreeze_blocks": args.unfreeze_blocks, | |
| "pretrained": not args.from_scratch}, | |
| args.out / "best.pt", | |
| ) | |
| print(f"best val score {best:.4f} -> {args.out / 'best.pt'}") | |
| if __name__ == "__main__": | |
| main() | |