Add distributed server implementation and protocol definitions
- Introduced `protocol.py` for shared data models used in server/client communication, including request and response schemas for registration, job submission, and status updates. - Implemented `server.py` to manage a SQLite job queue and client registry, handling job polling, status updates, and job completion. - Created a cheat sheet for server usage, detailing commands for starting the server, submitting jobs, and monitoring clients. - Added several experiment configuration files for various training setups, including geometry vector injections and baseline ensembles.
This commit is contained in:
@@ -15,3 +15,6 @@ scripts/deprecated/
|
|||||||
v3/results/*
|
v3/results/*
|
||||||
scripts/utility/backup_mirror_with_archive.sh
|
scripts/utility/backup_mirror_with_archive.sh
|
||||||
v3/distributed/logs/*
|
v3/distributed/logs/*
|
||||||
|
v4/configs/**/
|
||||||
|
v4/distributed/logs/*
|
||||||
|
v4/results/*
|
||||||
+100
-19
@@ -17,6 +17,7 @@ SegMapDataset — Dataset yielding (seg_tensor, label) pairs
|
|||||||
SegCNN — pretrained CNN adapted for segmentation-map input
|
SegCNN — pretrained CNN adapted for segmentation-map input
|
||||||
GeometryTower — TowerBase implementation (the main class to use)
|
GeometryTower — TowerBase implementation (the main class to use)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
@@ -35,15 +36,18 @@ from tqdm import tqdm
|
|||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from v3.classes.towerbase import TowerBase
|
from v3.classes.towerbase import TowerBase
|
||||||
|
from v3.classes.towerbase import TowerBase, EarlyPassContext
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Data record
|
# Data record
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class SegMapRecord:
|
class SegMapRecord:
|
||||||
"""One labelled eye sample for the seg-map CNN."""
|
"""One labelled eye sample for the seg-map CNN."""
|
||||||
|
|
||||||
sample_id: str
|
sample_id: str
|
||||||
image_path: Path # original fundus image (used by unet mode)
|
image_path: Path # original fundus image (used by unet mode)
|
||||||
annotation_disc: Path # contour (.txt) or mask (.bmp/.png)
|
annotation_disc: Path # contour (.txt) or mask (.bmp/.png)
|
||||||
@@ -59,6 +63,7 @@ class SegMapRecord:
|
|||||||
# Seg-map utilities
|
# Seg-map utilities
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _combine_masks(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray:
|
def _combine_masks(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray:
|
||||||
"""Combine binary disc and cup masks into a 3-class label map.
|
"""Combine binary disc and cup masks into a 3-class label map.
|
||||||
|
|
||||||
@@ -69,7 +74,7 @@ def _combine_masks(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray:
|
|||||||
"""
|
"""
|
||||||
disc = (disc_mask > 0).astype(np.uint8)
|
disc = (disc_mask > 0).astype(np.uint8)
|
||||||
cup = (cup_mask > 0).astype(np.uint8)
|
cup = (cup_mask > 0).astype(np.uint8)
|
||||||
cup = (cup & disc) # structural prior: cup must be inside disc
|
cup = cup & disc # structural prior: cup must be inside disc
|
||||||
seg = disc + cup # 0, 1 (rim), or 2 (cup)
|
seg = disc + cup # 0, 1 (rim), or 2 (cup)
|
||||||
return seg.astype(np.uint8)
|
return seg.astype(np.uint8)
|
||||||
|
|
||||||
@@ -86,7 +91,7 @@ def crop_to_disc(seg_map: np.ndarray) -> np.ndarray:
|
|||||||
return seg_map
|
return seg_map
|
||||||
r0, r1 = int(np.argmax(rows)), int(len(rows) - 1 - np.argmax(rows[::-1]))
|
r0, r1 = int(np.argmax(rows)), int(len(rows) - 1 - np.argmax(rows[::-1]))
|
||||||
c0, c1 = int(np.argmax(cols)), int(len(cols) - 1 - np.argmax(cols[::-1]))
|
c0, c1 = int(np.argmax(cols)), int(len(cols) - 1 - np.argmax(cols[::-1]))
|
||||||
return seg_map[r0:r1 + 1, c0:c1 + 1]
|
return seg_map[r0 : r1 + 1, c0 : c1 + 1]
|
||||||
|
|
||||||
|
|
||||||
def seg_map_to_tensor(
|
def seg_map_to_tensor(
|
||||||
@@ -120,12 +125,15 @@ def seg_map_to_tensor(
|
|||||||
# GT mask loading (pure NumPy / PIL — no CUDA, safe in DataLoader workers)
|
# GT mask loading (pure NumPy / PIL — no CUDA, safe in DataLoader workers)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _load_contour(path: Path) -> np.ndarray:
|
def _load_contour(path: Path) -> np.ndarray:
|
||||||
"""Load x,y contour pairs from a whitespace- or comma-delimited text file."""
|
"""Load x,y contour pairs from a whitespace- or comma-delimited text file."""
|
||||||
arr = np.zeros((0, 2), dtype=np.float32)
|
arr = np.zeros((0, 2), dtype=np.float32)
|
||||||
for delimiter in (",", None):
|
for delimiter in (",", None):
|
||||||
try:
|
try:
|
||||||
candidate = np.loadtxt(str(path), delimiter=delimiter, comments="#", dtype=np.float32)
|
candidate = np.loadtxt(
|
||||||
|
str(path), delimiter=delimiter, comments="#", dtype=np.float32
|
||||||
|
)
|
||||||
if candidate.size > 0:
|
if candidate.size > 0:
|
||||||
arr = candidate
|
arr = candidate
|
||||||
break
|
break
|
||||||
@@ -169,7 +177,9 @@ def _extract_masks_from_image(
|
|||||||
|
|
||||||
if arr.ndim == 2:
|
if arr.ndim == 2:
|
||||||
edges = np.concatenate([arr[0], arr[-1], arr[:, 0], arr[:, -1]])
|
edges = np.concatenate([arr[0], arr[-1], arr[:, 0], arr[:, -1]])
|
||||||
bg_val = int(np.argmax(np.bincount(edges.astype(np.int64).clip(0, 255), minlength=256)))
|
bg_val = int(
|
||||||
|
np.argmax(np.bincount(edges.astype(np.int64).clip(0, 255), minlength=256))
|
||||||
|
)
|
||||||
disc_arr = (arr != bg_val).astype(np.uint8)
|
disc_arr = (arr != bg_val).astype(np.uint8)
|
||||||
vals = np.unique(arr)
|
vals = np.unique(arr)
|
||||||
non_bg = vals[vals != bg_val]
|
non_bg = vals[vals != bg_val]
|
||||||
@@ -184,7 +194,9 @@ def _extract_masks_from_image(
|
|||||||
arr = np.array(img_rgb)
|
arr = np.array(img_rgb)
|
||||||
h, w, c = arr.shape
|
h, w, c = arr.shape
|
||||||
edges_rgb = np.concatenate([arr[0], arr[-1], arr[:, 0], arr[:, -1]], axis=0)
|
edges_rgb = np.concatenate([arr[0], arr[-1], arr[:, 0], arr[:, -1]], axis=0)
|
||||||
edge_colors, edge_counts = np.unique(edges_rgb.reshape(-1, c), axis=0, return_counts=True)
|
edge_colors, edge_counts = np.unique(
|
||||||
|
edges_rgb.reshape(-1, c), axis=0, return_counts=True
|
||||||
|
)
|
||||||
bg_color = edge_colors[int(np.argmax(edge_counts))]
|
bg_color = edge_colors[int(np.argmax(edge_counts))]
|
||||||
colors, counts = np.unique(arr.reshape(-1, c), axis=0, return_counts=True)
|
colors, counts = np.unique(arr.reshape(-1, c), axis=0, return_counts=True)
|
||||||
not_bg = np.any(colors != bg_color.reshape(1, -1), axis=1)
|
not_bg = np.any(colors != bg_color.reshape(1, -1), axis=1)
|
||||||
@@ -221,7 +233,9 @@ def load_gt_masks(rec: SegMapRecord, target_size: int) -> Tuple[np.ndarray, np.n
|
|||||||
|
|
||||||
# ---- Disc ----
|
# ---- Disc ----
|
||||||
if rec.annotation_type_disc == "mask":
|
if rec.annotation_type_disc == "mask":
|
||||||
disc_mask, cup_from_disc = _extract_masks_from_image(rec.annotation_disc, target_size)
|
disc_mask, cup_from_disc = _extract_masks_from_image(
|
||||||
|
rec.annotation_disc, target_size
|
||||||
|
)
|
||||||
if cup_from_disc.any():
|
if cup_from_disc.any():
|
||||||
cup_mask = cup_from_disc
|
cup_mask = cup_from_disc
|
||||||
else: # contour
|
else: # contour
|
||||||
@@ -250,6 +264,7 @@ def load_gt_masks(rec: SegMapRecord, target_size: int) -> Tuple[np.ndarray, np.n
|
|||||||
# U-Net fine-tuning dataset
|
# U-Net fine-tuning dataset
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class UNetFineTuneDataset(Dataset):
|
class UNetFineTuneDataset(Dataset):
|
||||||
"""Loads (image_tensor, mask_tensor) pairs for fine-tuning the U-Net."""
|
"""Loads (image_tensor, mask_tensor) pairs for fine-tuning the U-Net."""
|
||||||
|
|
||||||
@@ -294,6 +309,7 @@ class UNetFineTuneDataset(Dataset):
|
|||||||
# U-Net precomputation
|
# U-Net precomputation
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def precompute_unet_seg_maps(
|
def precompute_unet_seg_maps(
|
||||||
records: List[SegMapRecord],
|
records: List[SegMapRecord],
|
||||||
segmenter,
|
segmenter,
|
||||||
@@ -313,7 +329,7 @@ def precompute_unet_seg_maps(
|
|||||||
probs = torch.sigmoid(logits)[0].cpu().numpy()
|
probs = torch.sigmoid(logits)[0].cpu().numpy()
|
||||||
disc = (probs[0] > threshold).astype(np.uint8)
|
disc = (probs[0] > threshold).astype(np.uint8)
|
||||||
cup = (probs[1] > threshold).astype(np.uint8)
|
cup = (probs[1] > threshold).astype(np.uint8)
|
||||||
cup = (cup & disc)
|
cup = cup & disc
|
||||||
seg_maps.append(_combine_masks(disc, cup.astype(np.uint8)))
|
seg_maps.append(_combine_masks(disc, cup.astype(np.uint8)))
|
||||||
return seg_maps
|
return seg_maps
|
||||||
|
|
||||||
@@ -322,6 +338,7 @@ def precompute_unet_seg_maps(
|
|||||||
# SegMapDataset
|
# SegMapDataset
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class SegMapDataset(Dataset):
|
class SegMapDataset(Dataset):
|
||||||
"""PyTorch Dataset that yields (seg_tensor, label) pairs."""
|
"""PyTorch Dataset that yields (seg_tensor, label) pairs."""
|
||||||
|
|
||||||
@@ -445,7 +462,9 @@ class SegCNN(nn.Module):
|
|||||||
if pretrained:
|
if pretrained:
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
new_conv.weight.copy_(
|
new_conv.weight.copy_(
|
||||||
first_conv.weight.mean(dim=1, keepdim=True).expand_as(new_conv.weight)
|
first_conv.weight.mean(dim=1, keepdim=True).expand_as(
|
||||||
|
new_conv.weight
|
||||||
|
)
|
||||||
)
|
)
|
||||||
self._replace_first_conv(base, new_conv)
|
self._replace_first_conv(base, new_conv)
|
||||||
|
|
||||||
@@ -486,6 +505,7 @@ class SegCNN(nn.Module):
|
|||||||
# GeometryTower — TowerBase implementation
|
# GeometryTower — TowerBase implementation
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class GeometryTower(TowerBase, nn.Module):
|
class GeometryTower(TowerBase, nn.Module):
|
||||||
"""TowerBase implementation for the optic-disc/cup segmentation modality.
|
"""TowerBase implementation for the optic-disc/cup segmentation modality.
|
||||||
|
|
||||||
@@ -578,6 +598,48 @@ class GeometryTower(TowerBase, nn.Module):
|
|||||||
for p in self._seg_cnn.parameters():
|
for p in self._seg_cnn.parameters():
|
||||||
p.requires_grad = trainable
|
p.requires_grad = trainable
|
||||||
|
|
||||||
|
def early_pass(self, context: EarlyPassContext) -> None:
|
||||||
|
"""
|
||||||
|
V4 orchestrator hook: combines prepare_fold and augment_samples.
|
||||||
|
|
||||||
|
1. Pre-computes all segmentation maps for the current fold.
|
||||||
|
2. Injects them into the sample dicts held by the context object.
|
||||||
|
"""
|
||||||
|
if self._manifest_path is None:
|
||||||
|
raise ValueError("GeometryTower requires manifest_path")
|
||||||
|
|
||||||
|
all_paths: dict = {}
|
||||||
|
for split in (
|
||||||
|
context.eye_train,
|
||||||
|
context.bilat_train,
|
||||||
|
context.bilat_val,
|
||||||
|
context.bilat_test,
|
||||||
|
):
|
||||||
|
for s in split:
|
||||||
|
for slot in ("image_1", "image_2"):
|
||||||
|
p = s.get(slot)
|
||||||
|
if p is not None:
|
||||||
|
all_paths[str(Path(p).resolve())] = None
|
||||||
|
|
||||||
|
# 1. Pre-compute seg maps (from prepare_fold)
|
||||||
|
if self._geometry_source == "unet":
|
||||||
|
self._prepare_fold_unet(
|
||||||
|
list(all_paths.keys()), context.eye_train, context.device
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self._prepare_fold_gt(list(all_paths.keys()))
|
||||||
|
|
||||||
|
# 2. Inject into samples (from augment_samples)
|
||||||
|
for samples in (
|
||||||
|
context.eye_train,
|
||||||
|
context.bilat_train,
|
||||||
|
context.bilat_val,
|
||||||
|
context.bilat_test,
|
||||||
|
):
|
||||||
|
if samples:
|
||||||
|
# This modifies the list of dicts in the context object in-place
|
||||||
|
self.augment_samples(samples)
|
||||||
|
|
||||||
def prepare_fold(
|
def prepare_fold(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -617,7 +679,10 @@ class GeometryTower(TowerBase, nn.Module):
|
|||||||
(self._in_channels, self._target_size, self._target_size), dtype=np.float32
|
(self._in_channels, self._target_size, self._target_size), dtype=np.float32
|
||||||
)
|
)
|
||||||
for s in samples:
|
for s in samples:
|
||||||
for img_slot, seg_slot in (("image_1", "seg_map_1"), ("image_2", "seg_map_2")):
|
for img_slot, seg_slot in (
|
||||||
|
("image_1", "seg_map_1"),
|
||||||
|
("image_2", "seg_map_2"),
|
||||||
|
):
|
||||||
img_path = s.get(img_slot)
|
img_path = s.get(img_slot)
|
||||||
if img_path is None:
|
if img_path is None:
|
||||||
continue
|
continue
|
||||||
@@ -687,7 +752,10 @@ class GeometryTower(TowerBase, nn.Module):
|
|||||||
n_ok += 1
|
n_ok += 1
|
||||||
except Exception:
|
except Exception:
|
||||||
self._seg_cache[img_path] = self._seg_map_to_array(blank)
|
self._seg_cache[img_path] = self._seg_map_to_array(blank)
|
||||||
print(f"[GeometryTower] {n_ok}/{len(image_paths)} GT seg maps computed", flush=True)
|
print(
|
||||||
|
f"[GeometryTower] {n_ok}/{len(image_paths)} GT seg maps computed",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
def _prepare_fold_unet(self, image_paths: list, eye_train: list, device) -> None:
|
def _prepare_fold_unet(self, image_paths: list, eye_train: list, device) -> None:
|
||||||
"""Pre-compute U-Net seg maps, with optional per-fold fine-tuning."""
|
"""Pre-compute U-Net seg maps, with optional per-fold fine-tuning."""
|
||||||
@@ -716,14 +784,20 @@ class GeometryTower(TowerBase, nn.Module):
|
|||||||
target_size=segmenter.target_size,
|
target_size=segmenter.target_size,
|
||||||
normalize=self._unet_normalize,
|
normalize=self._unet_normalize,
|
||||||
),
|
),
|
||||||
batch_size=4, shuffle=True, num_workers=0,
|
batch_size=4,
|
||||||
|
shuffle=True,
|
||||||
|
num_workers=0,
|
||||||
|
)
|
||||||
|
optimizer = torch.optim.Adam(
|
||||||
|
segmenter.model.parameters(), lr=self._finetune_unet_lr
|
||||||
)
|
)
|
||||||
optimizer = torch.optim.Adam(segmenter.model.parameters(), lr=self._finetune_unet_lr)
|
|
||||||
criterion = torch.nn.BCEWithLogitsLoss()
|
criterion = torch.nn.BCEWithLogitsLoss()
|
||||||
segmenter.model.train()
|
segmenter.model.train()
|
||||||
for _ in range(self._finetune_unet_epochs):
|
for _ in range(self._finetune_unet_epochs):
|
||||||
for images, masks in ft_loader:
|
for images, masks in ft_loader:
|
||||||
images, masks = images.to(segmenter.device), masks.to(segmenter.device)
|
images, masks = images.to(segmenter.device), masks.to(
|
||||||
|
segmenter.device
|
||||||
|
)
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
criterion(segmenter.model(images), masks).backward()
|
criterion(segmenter.model(images), masks).backward()
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
@@ -735,10 +809,15 @@ class GeometryTower(TowerBase, nn.Module):
|
|||||||
)
|
)
|
||||||
records = [
|
records = [
|
||||||
SegMapRecord(
|
SegMapRecord(
|
||||||
sample_id="", image_path=Path(p),
|
sample_id="",
|
||||||
annotation_disc=Path(p), annotation_cup=Path(p),
|
image_path=Path(p),
|
||||||
annotation_type_disc="", annotation_type_cup="",
|
annotation_disc=Path(p),
|
||||||
patient_id=0, eye="", label=0,
|
annotation_cup=Path(p),
|
||||||
|
annotation_type_disc="",
|
||||||
|
annotation_type_cup="",
|
||||||
|
patient_id=0,
|
||||||
|
eye="",
|
||||||
|
label=0,
|
||||||
)
|
)
|
||||||
for p in image_paths
|
for p in image_paths
|
||||||
]
|
]
|
||||||
@@ -764,7 +843,8 @@ class GeometryTower(TowerBase, nn.Module):
|
|||||||
entry = manifest_index.get(key)
|
entry = manifest_index.get(key)
|
||||||
if entry is None:
|
if entry is None:
|
||||||
continue
|
continue
|
||||||
records.append(SegMapRecord(
|
records.append(
|
||||||
|
SegMapRecord(
|
||||||
sample_id="",
|
sample_id="",
|
||||||
image_path=Path(p),
|
image_path=Path(p),
|
||||||
annotation_disc=Path(entry["annotation_disc"]),
|
annotation_disc=Path(entry["annotation_disc"]),
|
||||||
@@ -774,5 +854,6 @@ class GeometryTower(TowerBase, nn.Module):
|
|||||||
patient_id=int(s.get("patient_id", 0)),
|
patient_id=int(s.get("patient_id", 0)),
|
||||||
eye=str(s.get("eye", "")),
|
eye=str(s.get("eye", "")),
|
||||||
label=int(s.get("label", 0)),
|
label=int(s.get("label", 0)),
|
||||||
))
|
)
|
||||||
|
)
|
||||||
return records
|
return records
|
||||||
|
|||||||
@@ -98,6 +98,10 @@ class FoldResult:
|
|||||||
fused_val_n: int = 0
|
fused_val_n: int = 0
|
||||||
fused_test_auc: float = float("nan")
|
fused_test_auc: float = float("nan")
|
||||||
fused_test_acc: float = float("nan")
|
fused_test_acc: float = float("nan")
|
||||||
|
fused_test_kappa: float = float("nan")
|
||||||
|
fused_test_f1: float = float("nan")
|
||||||
|
fused_test_ece: float = float("nan")
|
||||||
|
fused_test_n: int = 0
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -136,3 +140,5 @@ class FoldArtifacts:
|
|||||||
probs_test: Optional[np.ndarray] = None
|
probs_test: Optional[np.ndarray] = None
|
||||||
probs_test_img: Optional[np.ndarray] = None
|
probs_test_img: Optional[np.ndarray] = None
|
||||||
probs_test_md: Optional[np.ndarray] = None
|
probs_test_md: Optional[np.ndarray] = None
|
||||||
|
y_true_fused_test: Optional[np.ndarray] = None
|
||||||
|
probs_fused_test: Optional[np.ndarray] = None
|
||||||
|
|||||||
@@ -467,6 +467,10 @@ class V3HyperTower:
|
|||||||
np.save(fold_dir / "test_probs_img.npy", artifacts.probs_test_img)
|
np.save(fold_dir / "test_probs_img.npy", artifacts.probs_test_img)
|
||||||
if artifacts.probs_test_md is not None:
|
if artifacts.probs_test_md is not None:
|
||||||
np.save(fold_dir / "test_probs_cd.npy", artifacts.probs_test_md)
|
np.save(fold_dir / "test_probs_cd.npy", artifacts.probs_test_md)
|
||||||
|
if artifacts.y_true_fused_test is not None:
|
||||||
|
np.save(fold_dir / "test_y_true_fused.npy", artifacts.y_true_fused_test)
|
||||||
|
if artifacts.probs_fused_test is not None:
|
||||||
|
np.save(fold_dir / "test_probs_fused_head.npy", artifacts.probs_fused_test)
|
||||||
|
|
||||||
if pred_store is not None:
|
if pred_store is not None:
|
||||||
pred_store.save(tm_dir / "predictions.npz")
|
pred_store.save(tm_dir / "predictions.npz")
|
||||||
@@ -489,14 +493,17 @@ class V3HyperTower:
|
|||||||
if tower_mode in ("single", "classic"):
|
if tower_mode in ("single", "classic"):
|
||||||
_test_key = "classic_test"
|
_test_key = "classic_test"
|
||||||
elif tower_mode == "ensemble":
|
elif tower_mode == "ensemble":
|
||||||
_test_key = "ensemble_test"
|
_test_key = "fused_test" if fused_head else "ensemble_test"
|
||||||
elif tower_mode in ("bilateral", "siamese"):
|
elif tower_mode in ("bilateral", "siamese"):
|
||||||
_test_key = "bilat_test"
|
_test_key = "bilat_test"
|
||||||
else:
|
else:
|
||||||
_test_key = "classic_test"
|
_test_key = "classic_test"
|
||||||
|
mode_summary = {_test_key: summary.get(_test_key, {})}
|
||||||
|
if fused_head and tower_mode == "ensemble":
|
||||||
|
mode_summary["ensemble_test"] = summary.get("ensemble_test", {})
|
||||||
|
mode_summary["fused_best_val"] = summary.get("fused_best_val", {})
|
||||||
with (tm_dir / "summary.json").open("w") as f:
|
with (tm_dir / "summary.json").open("w") as f:
|
||||||
json.dump({"mode_summary": {_test_key: summary.get(_test_key, {})}},
|
json.dump({"mode_summary": mode_summary}, f, indent=2, default=str)
|
||||||
f, indent=2, default=str)
|
|
||||||
|
|
||||||
return out_dir
|
return out_dir
|
||||||
|
|
||||||
@@ -1524,10 +1531,17 @@ class V3HyperTower:
|
|||||||
y_fu_best = p_fu_best = None
|
y_fu_best = p_fu_best = None
|
||||||
if run_fused and single is not None:
|
if run_fused and single is not None:
|
||||||
y_fu_best, p_fu_best = collect_probs_fused(fused, val_loader, device)
|
y_fu_best, p_fu_best = collect_probs_fused(fused, val_loader, device)
|
||||||
|
if y_fu_best is not None and y_fu_best.size:
|
||||||
|
fu_acc_best = float((p_fu_best.argmax(1) == y_fu_best).mean())
|
||||||
|
snap_fused, _, _, _ = _tune_and_snap(
|
||||||
|
y_fu_best, p_fu_best, fu_acc_best, num_classes, args, args.ece_bins
|
||||||
|
)
|
||||||
|
|
||||||
# Test set evaluation (once, never seen during training)
|
# Test set evaluation (once, never seen during training)
|
||||||
snap_test: dict = {}
|
snap_test: dict = {}
|
||||||
|
snap_fused_test: dict = {}
|
||||||
y_test_out = p_test_out = p_test_img_out = p_test_md_out = None
|
y_test_out = p_test_out = p_test_img_out = p_test_md_out = None
|
||||||
|
y_fused_test_out = p_fused_test_out = None
|
||||||
|
|
||||||
if test_loader is not None:
|
if test_loader is not None:
|
||||||
if run_single and tower_mode == "ensemble":
|
if run_single and tower_mode == "ensemble":
|
||||||
@@ -1564,6 +1578,33 @@ class V3HyperTower:
|
|||||||
{"fused": p_test_out, "img": p_test_img_out, "md": p_test_md_out},
|
{"fused": p_test_out, "img": p_test_img_out, "md": p_test_md_out},
|
||||||
suffix="_test",
|
suffix="_test",
|
||||||
)
|
)
|
||||||
|
if run_fused and single is not None:
|
||||||
|
y_fused_test_out, p_fused_test_out = collect_probs_fused(
|
||||||
|
fused, test_loader, device
|
||||||
|
)
|
||||||
|
if y_fused_test_out is not None and y_fused_test_out.size:
|
||||||
|
fused_test_acc_raw = float(
|
||||||
|
(p_fused_test_out.argmax(1) == y_fused_test_out).mean()
|
||||||
|
)
|
||||||
|
snap_fused_test, _, _, _ = _tune_and_snap(
|
||||||
|
y_fused_test_out, p_fused_test_out, fused_test_acc_raw,
|
||||||
|
num_classes, args, args.ece_bins
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f" [fold {fold+1}] FUSED_HEAD TEST "
|
||||||
|
f"auc={snap_fused_test.get('auc', nan):.2f} "
|
||||||
|
f"acc={snap_fused_test.get('acc', nan):.2f} "
|
||||||
|
f"kappa={snap_fused_test.get('kappa', nan):.2f} "
|
||||||
|
f"f1={snap_fused_test.get('macro_f1', nan):.2f} "
|
||||||
|
f"ece={snap_fused_test.get('ece', nan):.2f} "
|
||||||
|
f"n={snap_fused_test.get('n', 0)}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
_save_predictions_csv(
|
||||||
|
fold_dir, mode, y_fused_test_out,
|
||||||
|
{"fused_head": p_fused_test_out},
|
||||||
|
suffix="_fused_head_test",
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
print(f" [fold {fold+1}] WARNING: no test samples for this fold.", flush=True)
|
print(f" [fold {fold+1}] WARNING: no test samples for this fold.", flush=True)
|
||||||
|
|
||||||
@@ -1630,7 +1671,12 @@ class V3HyperTower:
|
|||||||
fused_val_threshold=snap_fused.get("threshold", nan),
|
fused_val_threshold=snap_fused.get("threshold", nan),
|
||||||
fused_val_bias=_svf(snap_fused.get("bias")),
|
fused_val_bias=_svf(snap_fused.get("bias")),
|
||||||
fused_val_n=snap_fused.get("n", 0),
|
fused_val_n=snap_fused.get("n", 0),
|
||||||
fused_test_auc=nan, fused_test_acc=nan,
|
fused_test_auc=snap_fused_test.get("auc", nan),
|
||||||
|
fused_test_acc=snap_fused_test.get("acc", nan),
|
||||||
|
fused_test_kappa=snap_fused_test.get("kappa", nan),
|
||||||
|
fused_test_f1=snap_fused_test.get("macro_f1", nan),
|
||||||
|
fused_test_ece=snap_fused_test.get("ece", nan),
|
||||||
|
fused_test_n=snap_fused_test.get("n", 0),
|
||||||
), FoldArtifacts(
|
), FoldArtifacts(
|
||||||
y_true_classic=y_cl_best, probs_classic=p_cl_best,
|
y_true_classic=y_cl_best, probs_classic=p_cl_best,
|
||||||
y_true_ensemble=y_en_best, probs_ensemble=p_en_best,
|
y_true_ensemble=y_en_best, probs_ensemble=p_en_best,
|
||||||
@@ -1657,6 +1703,8 @@ class V3HyperTower:
|
|||||||
probs_test=p_test_out,
|
probs_test=p_test_out,
|
||||||
probs_test_img=p_test_img_out,
|
probs_test_img=p_test_img_out,
|
||||||
probs_test_md=p_test_md_out,
|
probs_test_md=p_test_md_out,
|
||||||
|
y_true_fused_test=y_fused_test_out,
|
||||||
|
probs_fused_test=p_fused_test_out,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -1685,6 +1733,7 @@ class V3HyperTower:
|
|||||||
("ensemble_test", "ensemble_test"),
|
("ensemble_test", "ensemble_test"),
|
||||||
("classic_test", "classic_test"),
|
("classic_test", "classic_test"),
|
||||||
("bilat_test", "bilat_test"),
|
("bilat_test", "bilat_test"),
|
||||||
|
("fused_test", "fused_test"),
|
||||||
]:
|
]:
|
||||||
sub = {}
|
sub = {}
|
||||||
for m in ["auc", "acc", "kappa", "f1", "ece"]:
|
for m in ["auc", "acc", "kappa", "f1", "ece"]:
|
||||||
|
|||||||
@@ -109,6 +109,12 @@ BACKBONES: Dict[str, BackboneSpec] = {
|
|||||||
strip=_strip_efficientnet,
|
strip=_strip_efficientnet,
|
||||||
blocks=_blocks_efficientnet,
|
blocks=_blocks_efficientnet,
|
||||||
),
|
),
|
||||||
|
"resnet18": BackboneSpec(
|
||||||
|
ctor=models.resnet18,
|
||||||
|
weights_default=models.ResNet18_Weights.DEFAULT,
|
||||||
|
strip=_strip_resnet,
|
||||||
|
blocks=_blocks_resnet,
|
||||||
|
),
|
||||||
"resnet50": BackboneSpec(
|
"resnet50": BackboneSpec(
|
||||||
ctor=models.resnet50,
|
ctor=models.resnet50,
|
||||||
weights_default=models.ResNet50_Weights.DEFAULT,
|
weights_default=models.ResNet50_Weights.DEFAULT,
|
||||||
|
|||||||
@@ -44,6 +44,37 @@ class ImageTransformConfig:
|
|||||||
]
|
]
|
||||||
return transforms.Compose(ops)
|
return transforms.Compose(ops)
|
||||||
|
|
||||||
|
def build_precache(self) -> transforms.Compose:
|
||||||
|
"""Deterministic prefix: PIL → resized CHW float32 in [0, 1].
|
||||||
|
|
||||||
|
Output is suitable for caching; per-batch ``build_postcache`` finishes
|
||||||
|
the pipeline (augment + normalize) on tensors.
|
||||||
|
"""
|
||||||
|
return transforms.Compose([
|
||||||
|
transforms.Resize(self.resize_size),
|
||||||
|
transforms.CenterCrop(self.crop_size),
|
||||||
|
transforms.ToTensor(),
|
||||||
|
])
|
||||||
|
|
||||||
|
def build_postcache(self) -> transforms.Compose:
|
||||||
|
"""Per-batch tail run on cached float32 [0, 1] CHW tensors.
|
||||||
|
|
||||||
|
Augmentations operate on tensors (torchvision v1 supports this for
|
||||||
|
Flip/Rotation/ColorJitter on tensor input). Normalize is applied last.
|
||||||
|
"""
|
||||||
|
ops = []
|
||||||
|
if self.augment:
|
||||||
|
if self.hflip:
|
||||||
|
ops.append(transforms.RandomHorizontalFlip())
|
||||||
|
if self.vflip:
|
||||||
|
ops.append(transforms.RandomVerticalFlip())
|
||||||
|
if self.rotation_deg:
|
||||||
|
ops.append(transforms.RandomRotation(self.rotation_deg))
|
||||||
|
if self.color_jitter:
|
||||||
|
ops.append(transforms.ColorJitter(*self.color_jitter))
|
||||||
|
ops.append(transforms.Normalize(mean=self.mean, std=self.std))
|
||||||
|
return transforms.Compose(ops)
|
||||||
|
|
||||||
|
|
||||||
def backbone_transform_config(backbone_name: str, augment: bool = True) -> ImageTransformConfig:
|
def backbone_transform_config(backbone_name: str, augment: bool = True) -> ImageTransformConfig:
|
||||||
"""Build an ImageTransformConfig using the backbone's default normalisation stats."""
|
"""Build an ImageTransformConfig using the backbone's default normalisation stats."""
|
||||||
@@ -64,3 +95,15 @@ def build_backbone_transform(backbone_name: str, augment: bool = True) -> transf
|
|||||||
def build_eval_transform(backbone_name: str) -> transforms.Compose:
|
def build_eval_transform(backbone_name: str) -> transforms.Compose:
|
||||||
"""Deterministic eval transform — no augmentation, backbone-matched normalisation."""
|
"""Deterministic eval transform — no augmentation, backbone-matched normalisation."""
|
||||||
return build_backbone_transform(backbone_name, augment=False)
|
return build_backbone_transform(backbone_name, augment=False)
|
||||||
|
|
||||||
|
|
||||||
|
def build_split_transforms(
|
||||||
|
backbone_name: str, augment: bool = True
|
||||||
|
) -> tuple[transforms.Compose, transforms.Compose]:
|
||||||
|
"""Return (precache, postcache) transform pair for tensor-cached image towers.
|
||||||
|
|
||||||
|
precache : PIL → CHW float32 in [0, 1] (deterministic, run once at fill)
|
||||||
|
postcache : tensor → augmented + normalized tensor (run per batch)
|
||||||
|
"""
|
||||||
|
cfg = backbone_transform_config(backbone_name, augment=augment)
|
||||||
|
return cfg.build_precache(), cfg.build_postcache()
|
||||||
|
|||||||
@@ -0,0 +1,212 @@
|
|||||||
|
"""unet — REFUGE-trained UNet wrapper for v4.
|
||||||
|
|
||||||
|
Lean accessory module: model definition + a thin segmenter wrapper that handles
|
||||||
|
weight loading, preprocessing, inference, and fine-tuning.
|
||||||
|
|
||||||
|
Used by:
|
||||||
|
- GeometrySegEncoder tower (produces disc/cup seg maps as CNN input)
|
||||||
|
- (future) ImageEncoder cropping (locates disc bbox for image cropping)
|
||||||
|
|
||||||
|
The segmenter is intentionally domain-agnostic: it takes PIL images in and
|
||||||
|
returns binary (disc, cup) numpy masks. Fine-tuning consumes any DataLoader
|
||||||
|
yielding (image_tensor, mask_tensor) pairs — mask preparation (parsing GT
|
||||||
|
contour files, etc.) lives in the consumer.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image
|
||||||
|
from PIL.Image import Resampling
|
||||||
|
from torch import nn
|
||||||
|
from torchvision import transforms
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Model
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class UNet(nn.Module):
|
||||||
|
def __init__(self, in_channels: int = 3, base_channels: int = 32, out_channels: int = 2):
|
||||||
|
super().__init__()
|
||||||
|
self.enc1 = self._block(in_channels, base_channels)
|
||||||
|
self.enc2 = self._block(base_channels, base_channels * 2)
|
||||||
|
self.enc3 = self._block(base_channels * 2, base_channels * 4)
|
||||||
|
self.enc4 = self._block(base_channels * 4, base_channels * 8)
|
||||||
|
|
||||||
|
self.pool = nn.MaxPool2d(2)
|
||||||
|
self.bottleneck = self._block(base_channels * 8, base_channels * 16)
|
||||||
|
|
||||||
|
self.up4 = nn.ConvTranspose2d(base_channels * 16, base_channels * 8, 2, stride=2)
|
||||||
|
self.dec4 = self._block(base_channels * 16, base_channels * 8)
|
||||||
|
self.up3 = nn.ConvTranspose2d(base_channels * 8, base_channels * 4, 2, stride=2)
|
||||||
|
self.dec3 = self._block(base_channels * 8, base_channels * 4)
|
||||||
|
self.up2 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, 2, stride=2)
|
||||||
|
self.dec2 = self._block(base_channels * 4, base_channels * 2)
|
||||||
|
self.up1 = nn.ConvTranspose2d(base_channels * 2, base_channels, 2, stride=2)
|
||||||
|
self.dec1 = self._block(base_channels * 2, base_channels)
|
||||||
|
|
||||||
|
self.out_conv = nn.Conv2d(base_channels, out_channels, kernel_size=1)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _block(in_ch: int, out_ch: int) -> nn.Module:
|
||||||
|
return nn.Sequential(
|
||||||
|
nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),
|
||||||
|
nn.BatchNorm2d(out_ch),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),
|
||||||
|
nn.BatchNorm2d(out_ch),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
e1 = self.enc1(x)
|
||||||
|
e2 = self.enc2(self.pool(e1))
|
||||||
|
e3 = self.enc3(self.pool(e2))
|
||||||
|
e4 = self.enc4(self.pool(e3))
|
||||||
|
b = self.bottleneck(self.pool(e4))
|
||||||
|
|
||||||
|
d4 = self.dec4(torch.cat([self.up4(b), e4], dim=1))
|
||||||
|
d3 = self.dec3(torch.cat([self.up3(d4), e3], dim=1))
|
||||||
|
d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1))
|
||||||
|
d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))
|
||||||
|
return self.out_conv(d1)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Segmenter wrapper
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
_IMAGENET_MEAN = (0.485, 0.456, 0.406)
|
||||||
|
_IMAGENET_STD = (0.229, 0.224, 0.225)
|
||||||
|
|
||||||
|
|
||||||
|
class UNetSegmenter:
|
||||||
|
"""Wraps a UNet with preprocessing, weight loading, inference, and fine-tuning.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
target_size : square resolution UNet operates at (default 512)
|
||||||
|
normalize : "per_image" | "imagenet" | "none"
|
||||||
|
device : torch device string; defaults to cuda if available
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
target_size: int = 512,
|
||||||
|
normalize: str = "per_image",
|
||||||
|
device: str | torch.device | None = None,
|
||||||
|
in_channels: int = 3,
|
||||||
|
base_channels: int = 32,
|
||||||
|
out_channels: int = 2,
|
||||||
|
):
|
||||||
|
self.target_size = target_size
|
||||||
|
self.normalize = normalize
|
||||||
|
self.device = (
|
||||||
|
torch.device(device) if device is not None
|
||||||
|
else torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
)
|
||||||
|
self.model = UNet(in_channels, base_channels, out_channels).to(self.device)
|
||||||
|
self._to_tensor = transforms.ToTensor()
|
||||||
|
|
||||||
|
# ── lifecycle ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def to(self, device: str | torch.device) -> "UNetSegmenter":
|
||||||
|
self.device = torch.device(device)
|
||||||
|
self.model.to(self.device)
|
||||||
|
return self
|
||||||
|
|
||||||
|
def load_weights(self, path: str | Path) -> "UNetSegmenter":
|
||||||
|
"""Load a UNet checkpoint (raw state_dict or {'model': state_dict})."""
|
||||||
|
state = torch.load(Path(path), map_location=self.device, weights_only=False)
|
||||||
|
sd = state["model"] if isinstance(state, dict) and "model" in state else state
|
||||||
|
self.model.load_state_dict(sd)
|
||||||
|
self.model.eval()
|
||||||
|
return self
|
||||||
|
|
||||||
|
# ── preprocessing ────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _normalize_tensor(self, t: torch.Tensor) -> torch.Tensor:
|
||||||
|
if self.normalize == "per_image":
|
||||||
|
mean = t.mean(dim=(-2, -1), keepdim=True)
|
||||||
|
std = t.std (dim=(-2, -1), keepdim=True).clamp(min=1e-6)
|
||||||
|
return (t - mean) / std
|
||||||
|
if self.normalize == "imagenet":
|
||||||
|
mean = torch.tensor(_IMAGENET_MEAN, device=t.device).view(-1, 1, 1)
|
||||||
|
std = torch.tensor(_IMAGENET_STD, device=t.device).view(-1, 1, 1)
|
||||||
|
return (t - mean) / std
|
||||||
|
return t
|
||||||
|
|
||||||
|
def preprocess(self, image: Image.Image) -> torch.Tensor:
|
||||||
|
"""PIL image → normalized (C, H, W) tensor on segmenter device."""
|
||||||
|
resized = image.convert("RGB").resize(
|
||||||
|
(self.target_size, self.target_size), Resampling.BILINEAR
|
||||||
|
)
|
||||||
|
return self._normalize_tensor(self._to_tensor(resized).to(self.device))
|
||||||
|
|
||||||
|
# ── inference ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def predict(
|
||||||
|
self,
|
||||||
|
image: Image.Image,
|
||||||
|
*,
|
||||||
|
threshold: float = 0.5,
|
||||||
|
tta: bool = False,
|
||||||
|
) -> tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""Single image → (disc_mask, cup_mask) binary uint8 arrays at target_size.
|
||||||
|
|
||||||
|
cup_mask is restricted to disc area (cup ⊆ disc).
|
||||||
|
"""
|
||||||
|
self.model.eval()
|
||||||
|
x = self.preprocess(image).unsqueeze(0)
|
||||||
|
logits = self.model(x)
|
||||||
|
if tta:
|
||||||
|
log_h = torch.flip(self.model(torch.flip(x, dims=[3])), dims=[3])
|
||||||
|
log_v = torch.flip(self.model(torch.flip(x, dims=[2])), dims=[2])
|
||||||
|
logits = (logits + log_h + log_v) / 3.0
|
||||||
|
probs = torch.sigmoid(logits)[0].cpu().numpy()
|
||||||
|
disc = (probs[0] > threshold).astype(np.uint8)
|
||||||
|
cup = ((probs[1] > threshold) & (disc > 0)).astype(np.uint8)
|
||||||
|
return disc, cup
|
||||||
|
|
||||||
|
# ── fine-tuning ──────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def finetune(
|
||||||
|
self,
|
||||||
|
dataloader,
|
||||||
|
*,
|
||||||
|
epochs: int = 10,
|
||||||
|
lr: float = 1e-5,
|
||||||
|
log_prefix: str = "[UNetSegmenter]",
|
||||||
|
) -> "UNetSegmenter":
|
||||||
|
"""Fine-tune on (image_tensor, mask_tensor) pairs.
|
||||||
|
|
||||||
|
image_tensor : (B, C, H, W) — already preprocessed (normalized)
|
||||||
|
mask_tensor : (B, 2, H, W) float32 — channel 0 disc, channel 1 cup
|
||||||
|
"""
|
||||||
|
import time
|
||||||
|
opt = torch.optim.Adam(self.model.parameters(), lr=lr)
|
||||||
|
crit = nn.BCEWithLogitsLoss()
|
||||||
|
for ep in range(1, epochs + 1):
|
||||||
|
self.model.train()
|
||||||
|
running, n_batches, t0 = 0.0, 0, time.time()
|
||||||
|
for img, mask in dataloader:
|
||||||
|
img, mask = img.to(self.device), mask.to(self.device)
|
||||||
|
opt.zero_grad()
|
||||||
|
loss = crit(self.model(img), mask)
|
||||||
|
loss.backward()
|
||||||
|
opt.step()
|
||||||
|
running += float(loss.item())
|
||||||
|
n_batches += 1
|
||||||
|
avg = running / max(n_batches, 1)
|
||||||
|
print(
|
||||||
|
f" {log_prefix} ep{ep:03d}/{epochs:03d} loss={avg:.4f} "
|
||||||
|
f"({time.time() - t0:.1f}s)",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
self.model.eval()
|
||||||
|
return self
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
"""mono_bridge — MonoBridge: passthrough for single-tower fusion stages.
|
||||||
|
|
||||||
|
The v4 stage runner always expects tower → bridge → head. For configs that
|
||||||
|
have only one tower feeding a head, MonoBridge is the no-op bridge that lets
|
||||||
|
the architecture be "head sits directly on tower" without any extra projection,
|
||||||
|
SE, or fusion logic.
|
||||||
|
|
||||||
|
Optional LayerNorm is exposed for consistency with FusionBridge but defaults
|
||||||
|
off to keep the embedding numerically identical to the tower's output.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
|
||||||
|
class MonoBridge(nn.Module):
|
||||||
|
"""Single-input passthrough bridge.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
input_dims : list[int] — must be length 1
|
||||||
|
use_ln : if True, wrap the embedding in a LayerNorm
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, input_dims: list[int], use_ln: bool = False):
|
||||||
|
super().__init__()
|
||||||
|
if len(input_dims) != 1:
|
||||||
|
raise ValueError(
|
||||||
|
f"MonoBridge expects exactly 1 input dim, got {len(input_dims)}"
|
||||||
|
)
|
||||||
|
self.out_dim = input_dims[0]
|
||||||
|
self.ln = nn.LayerNorm(self.out_dim) if use_ln else nn.Identity()
|
||||||
|
|
||||||
|
def forward(self, embeddings: list[torch.Tensor]) -> torch.Tensor:
|
||||||
|
if len(embeddings) != 1:
|
||||||
|
raise ValueError(
|
||||||
|
f"MonoBridge forward expects 1 embedding, got {len(embeddings)}"
|
||||||
|
)
|
||||||
|
return self.ln(embeddings[0])
|
||||||
|
|
||||||
|
def set_phase(self, phase: str) -> None:
|
||||||
|
"""Freeze during tower_warmup; trainable otherwise (matches FusionBridge)."""
|
||||||
|
enabled = phase not in ("tower_warmup", "cd_warmup")
|
||||||
|
for p in self.parameters():
|
||||||
|
p.requires_grad_(enabled)
|
||||||
@@ -205,7 +205,7 @@ class PredictionStore:
|
|||||||
grp.create_dataset("y_true", data=buf.y_true)
|
grp.create_dataset("y_true", data=buf.y_true)
|
||||||
grp.create_dataset("loss", data=buf.loss)
|
grp.create_dataset("loss", data=buf.loss)
|
||||||
grp.create_dataset("head_names", data=np.array(buf.head_names, dtype=object), dtype=_STR_DT)
|
grp.create_dataset("head_names", data=np.array(buf.head_names, dtype=object), dtype=_STR_DT)
|
||||||
grp.create_dataset("split", data=buf.split.astype(str), dtype=_STR_DT)
|
grp.create_dataset("split", data=buf.split, dtype=_STR_DT)
|
||||||
_write_entity_ids(grp, buf.entity_ids)
|
_write_entity_ids(grp, buf.entity_ids)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -355,7 +355,7 @@ class FeatureStore:
|
|||||||
for phase, buf in self._phases.items():
|
for phase, buf in self._phases.items():
|
||||||
grp = f.create_group(phase)
|
grp = f.create_group(phase)
|
||||||
grp.create_dataset("y_true", data=buf.y_true)
|
grp.create_dataset("y_true", data=buf.y_true)
|
||||||
grp.create_dataset("split", data=buf.split.astype(str), dtype=_STR_DT)
|
grp.create_dataset("split", data=buf.split, dtype=_STR_DT)
|
||||||
_write_entity_ids(grp, buf.entity_ids)
|
_write_entity_ids(grp, buf.entity_ids)
|
||||||
for head, (arr, _) in buf._heads.items():
|
for head, (arr, _) in buf._heads.items():
|
||||||
grp.create_dataset(head, data=arr, compression="gzip", compression_opts=4)
|
grp.create_dataset(head, data=arr, compression="gzip", compression_opts=4)
|
||||||
|
|||||||
@@ -0,0 +1,574 @@
|
|||||||
|
"""fundus_images — disc/cup geometry for fundus image profiles.
|
||||||
|
|
||||||
|
Contains all fundus-specific geometry logic: mask parsing, feature computation,
|
||||||
|
and source-specific loaders. Profiles whose ImageDataView supports geometry
|
||||||
|
should implement build_geometry_loader(source, **kwargs) and/or
|
||||||
|
build_seg_map_loader(source, **kwargs) and delegate here.
|
||||||
|
|
||||||
|
Geometry vectors (5 scalar CDR features per eye)
|
||||||
|
-----------------------------------------------
|
||||||
|
build_geometry_loader("gt", contour_dir=...) → GTGeometryLoader
|
||||||
|
|
||||||
|
Seg maps (3-class disc/cup label map per eye, fed to a CNN tower)
|
||||||
|
-----------------------------------------------------------------
|
||||||
|
build_seg_map_loader("gt", contour_dir=..., ...) → GTSegMapLoader
|
||||||
|
build_seg_map_loader("unet", weights_path=..., ...) → UNetSegMapLoader
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import copy
|
||||||
|
from collections import Counter
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Iterable, Tuple
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image, ImageDraw
|
||||||
|
from PIL.Image import Resampling
|
||||||
|
from torch.utils.data import DataLoader, Dataset
|
||||||
|
|
||||||
|
EPS = 1e-6
|
||||||
|
|
||||||
|
_FEATURE_DIM = 5
|
||||||
|
_FEATURE_NAMES = ["area_cdr", "rim_ratio", "vertical_cdr", "horizontal_cdr", "centre_shift"]
|
||||||
|
|
||||||
|
_MASK_SIZE = (512, 512) # canonical rasterisation size; CDR ratios are scale-invariant
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Mask utilities
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def disc_cup_from_mask_image(mask_img: Image.Image) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""Return binary (disc, cup) masks from a REFUGE-style colour annotation image."""
|
||||||
|
arr = np.asarray(mask_img)
|
||||||
|
if arr.ndim == 3:
|
||||||
|
h, w, c = arr.shape
|
||||||
|
border = np.concatenate(
|
||||||
|
[arr[0, :, :], arr[-1, :, :], arr[:, 0, :], arr[:, -1, :]], axis=0
|
||||||
|
)
|
||||||
|
bg_color = Counter(map(tuple, border)).most_common(1)[0][0]
|
||||||
|
colors = Counter(map(tuple, arr.reshape(-1, c)))
|
||||||
|
colors.pop(bg_color, None)
|
||||||
|
disc = (~np.all(arr == bg_color, axis=-1)).astype(np.uint8)
|
||||||
|
if colors:
|
||||||
|
cup_color = min(colors.keys(), key=lambda col: sum(col))
|
||||||
|
cup = np.all(arr == cup_color, axis=-1).astype(np.uint8)
|
||||||
|
else:
|
||||||
|
cup = np.zeros((h, w), dtype=np.uint8)
|
||||||
|
else:
|
||||||
|
border = np.concatenate([arr[0, :], arr[-1, :], arr[:, 0], arr[:, -1]])
|
||||||
|
bg_value = Counter(border.tolist()).most_common(1)[0][0]
|
||||||
|
disc = (arr != bg_value).astype(np.uint8)
|
||||||
|
fg = arr[arr != bg_value]
|
||||||
|
cup = (arr == int(np.min(fg))).astype(np.uint8) if fg.size > 0 else np.zeros_like(arr)
|
||||||
|
cup = (cup > 0) & (disc > 0)
|
||||||
|
return disc.astype(np.uint8), cup.astype(np.uint8)
|
||||||
|
|
||||||
|
|
||||||
|
def _contour_to_mask(coords: np.ndarray, size: Tuple[int, int]) -> np.ndarray:
|
||||||
|
"""Rasterize a polygon contour (Nx2 xy array) into a binary mask of (width, height)."""
|
||||||
|
from PIL import ImageDraw
|
||||||
|
if coords is None or coords.size == 0:
|
||||||
|
return np.zeros((size[1], size[0]), dtype=np.uint8)
|
||||||
|
img = Image.new("L", size, 0)
|
||||||
|
draw = ImageDraw.Draw(img)
|
||||||
|
draw.polygon([tuple(map(float, pt)) for pt in coords], outline=1, fill=1)
|
||||||
|
return np.array(img, dtype=np.uint8)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Feature computation
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def compute_geometry_features(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray:
|
||||||
|
"""Compute 5 cup/disc structural descriptors from binary masks.
|
||||||
|
|
||||||
|
Returns float32 [area_cdr, rim_ratio, vertical_cdr, horizontal_cdr, centre_shift].
|
||||||
|
"""
|
||||||
|
disc = (disc_mask > 0).astype(np.float32)
|
||||||
|
cup = (cup_mask > 0).astype(np.float32)
|
||||||
|
|
||||||
|
disc_area = disc.sum()
|
||||||
|
cup_area = cup.sum()
|
||||||
|
area_cdr = cup_area / (disc_area + EPS)
|
||||||
|
rim_ratio = (disc_area - cup_area) / (disc_area + EPS)
|
||||||
|
|
||||||
|
disc_h = float(np.any(disc > 0, axis=1).sum())
|
||||||
|
cup_h = float(np.any(cup > 0, axis=1).sum())
|
||||||
|
disc_w = float(np.any(disc > 0, axis=0).sum())
|
||||||
|
cup_w = float(np.any(cup > 0, axis=0).sum())
|
||||||
|
|
||||||
|
vertical_cdr = cup_h / (disc_h + EPS)
|
||||||
|
horizontal_cdr = cup_w / (disc_w + EPS)
|
||||||
|
|
||||||
|
def _centre(m: np.ndarray) -> Tuple[float, float]:
|
||||||
|
coords = np.argwhere(m > 0)
|
||||||
|
if coords.size == 0:
|
||||||
|
return 0.5, 0.5
|
||||||
|
ys, xs = coords[:, 0], coords[:, 1]
|
||||||
|
return float(xs.mean()) / m.shape[1], float(ys.mean()) / m.shape[0]
|
||||||
|
|
||||||
|
dcx, dcy = _centre(disc)
|
||||||
|
ccx, ccy = _centre(cup)
|
||||||
|
centre_shift = float(np.hypot(ccx - dcx, ccy - dcy))
|
||||||
|
|
||||||
|
return np.array(
|
||||||
|
[area_cdr, rim_ratio, vertical_cdr, horizontal_cdr, centre_shift],
|
||||||
|
dtype=np.float32,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Loaders
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class GTGeometryLoader:
|
||||||
|
"""Pre-computes per-eye geometry vectors from PAPILA GT contour annotations.
|
||||||
|
|
||||||
|
File naming: RET{pid:03d}{eye}_{disc|cup}_exp{n}.txt
|
||||||
|
Averages exp1 and exp2 when both are present; zero vector for missing entries.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
loader = GTGeometryLoader(contour_dir)
|
||||||
|
loader.precompute(df, patient_col="Patient ID")
|
||||||
|
vecs = loader.all_vectors() # {(pid, eye): ndarray}
|
||||||
|
"""
|
||||||
|
|
||||||
|
_EXPERTS = (1, 2)
|
||||||
|
|
||||||
|
feature_dim = _FEATURE_DIM
|
||||||
|
feature_names = _FEATURE_NAMES
|
||||||
|
|
||||||
|
def __init__(self, contour_dir: str | Path) -> None:
|
||||||
|
self._dir = Path(contour_dir)
|
||||||
|
self._cache: dict[tuple, np.ndarray] = {}
|
||||||
|
|
||||||
|
def precompute(self, df, patient_col: str = "Patient ID") -> None:
|
||||||
|
n_ok = 0
|
||||||
|
for _, row in df.iterrows():
|
||||||
|
pid = int(row[patient_col])
|
||||||
|
eye = str(row.get("eyeID", "OD"))
|
||||||
|
key = (pid, eye)
|
||||||
|
if key in self._cache:
|
||||||
|
continue
|
||||||
|
vec = self._compute(pid, eye)
|
||||||
|
self._cache[key] = vec if vec is not None else np.zeros(self.feature_dim, dtype=np.float32)
|
||||||
|
if vec is not None:
|
||||||
|
n_ok += 1
|
||||||
|
print(f"[GTGeometryLoader] {n_ok}/{len(self._cache)} geometry vectors computed", flush=True)
|
||||||
|
|
||||||
|
def all_vectors(self) -> dict:
|
||||||
|
return dict(self._cache)
|
||||||
|
|
||||||
|
def _compute(self, pid: int, eye: str) -> "np.ndarray | None":
|
||||||
|
stem = f"RET{pid:03d}{eye}"
|
||||||
|
vecs: list[np.ndarray] = []
|
||||||
|
for exp in self._EXPERTS:
|
||||||
|
disc_path = self._dir / f"{stem}_disc_exp{exp}.txt"
|
||||||
|
cup_path = self._dir / f"{stem}_cup_exp{exp}.txt"
|
||||||
|
if not disc_path.exists():
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
disc_c = np.loadtxt(disc_path)
|
||||||
|
if disc_c.ndim == 1:
|
||||||
|
disc_c = disc_c.reshape(-1, 2)
|
||||||
|
disc_mask = _contour_to_mask(disc_c, _MASK_SIZE)
|
||||||
|
if cup_path.exists():
|
||||||
|
cup_c = np.loadtxt(cup_path)
|
||||||
|
if cup_c.ndim == 1:
|
||||||
|
cup_c = cup_c.reshape(-1, 2)
|
||||||
|
cup_mask = _contour_to_mask(cup_c, _MASK_SIZE)
|
||||||
|
else:
|
||||||
|
cup_mask = np.zeros((_MASK_SIZE[1], _MASK_SIZE[0]), dtype=np.uint8)
|
||||||
|
cup_mask = ((cup_mask > 0) & (disc_mask > 0)).astype(np.uint8)
|
||||||
|
vecs.append(compute_geometry_features(disc_mask, cup_mask))
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
if not vecs:
|
||||||
|
return None
|
||||||
|
return np.stack(vecs).mean(axis=0).astype(np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Seg-map utilities (shared by GT and UNet seg-map loaders)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _combine_disc_cup(disc: np.ndarray, cup: np.ndarray) -> np.ndarray:
|
||||||
|
"""Merge binary disc + cup masks into a uint8 label map: 0=bg, 1=rim, 2=cup."""
|
||||||
|
disc = (disc > 0).astype(np.uint8)
|
||||||
|
cup = ((cup > 0) & (disc > 0)).astype(np.uint8)
|
||||||
|
return (disc + cup).astype(np.uint8)
|
||||||
|
|
||||||
|
|
||||||
|
def _crop_to_disc_bbox(seg_map: np.ndarray) -> np.ndarray:
|
||||||
|
"""Crop a label map tightly to the disc bounding box (anywhere seg_map > 0)."""
|
||||||
|
rows = np.any(seg_map > 0, axis=1)
|
||||||
|
cols = np.any(seg_map > 0, axis=0)
|
||||||
|
if not rows.any():
|
||||||
|
return seg_map
|
||||||
|
r0, r1 = int(np.argmax(rows)), int(len(rows) - 1 - np.argmax(rows[::-1]))
|
||||||
|
c0, c1 = int(np.argmax(cols)), int(len(cols) - 1 - np.argmax(cols[::-1]))
|
||||||
|
return seg_map[r0:r1 + 1, c0:c1 + 1]
|
||||||
|
|
||||||
|
|
||||||
|
def _seg_map_to_array(
|
||||||
|
seg_map: np.ndarray, channels: int, target_size: int
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""Resize a {0,1,2} seg map and convert to a (C, H, W) float32 array.
|
||||||
|
|
||||||
|
channels=1 → (1, H, W) values in {0, 0.5, 1.0}
|
||||||
|
channels=3 → (3, H, W) one-hot [bg, rim, cup]
|
||||||
|
"""
|
||||||
|
pil = Image.fromarray(seg_map.astype(np.uint8), mode="L").resize(
|
||||||
|
(target_size, target_size), Resampling.NEAREST
|
||||||
|
)
|
||||||
|
arr = np.array(pil, dtype=np.uint8)
|
||||||
|
if channels == 1:
|
||||||
|
return (arr.astype(np.float32) / 2.0)[None, :, :]
|
||||||
|
if channels == 3:
|
||||||
|
return np.stack([
|
||||||
|
(arr == 0).astype(np.float32),
|
||||||
|
(arr == 1).astype(np.float32),
|
||||||
|
(arr == 2).astype(np.float32),
|
||||||
|
], axis=0)
|
||||||
|
raise ValueError(f"channels must be 1 or 3, got {channels}")
|
||||||
|
|
||||||
|
|
||||||
|
def _load_papila_contour(path: Path) -> np.ndarray:
|
||||||
|
"""Load (x, y) contour pairs from a PAPILA whitespace/comma-delimited text file."""
|
||||||
|
for delim in (",", None):
|
||||||
|
try:
|
||||||
|
arr = np.loadtxt(str(path), delimiter=delim, comments="#", dtype=np.float32)
|
||||||
|
if arr.size > 0 and arr.ndim >= 1:
|
||||||
|
if arr.ndim == 1:
|
||||||
|
arr = arr.reshape(-1, 2)
|
||||||
|
if arr.shape[1] >= 2:
|
||||||
|
return arr[:, :2]
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
return np.zeros((0, 2), dtype=np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def _papila_disc_cup_masks(
|
||||||
|
pid: int, eye: str, contour_dir: Path, image_size: Tuple[int, int],
|
||||||
|
mask_size: int, experts: Tuple[int, ...] = (1, 2),
|
||||||
|
) -> Tuple[np.ndarray, np.ndarray] | None:
|
||||||
|
"""Average masks across PAPILA experts. Returns (disc, cup) at mask_size, or None."""
|
||||||
|
stem = f"RET{pid:03d}{eye}"
|
||||||
|
discs, cups = [], []
|
||||||
|
for exp in experts:
|
||||||
|
disc_path = contour_dir / f"{stem}_disc_exp{exp}.txt"
|
||||||
|
cup_path = contour_dir / f"{stem}_cup_exp{exp}.txt"
|
||||||
|
if not disc_path.exists():
|
||||||
|
continue
|
||||||
|
disc_c = _load_papila_contour(disc_path)
|
||||||
|
if len(disc_c) < 3:
|
||||||
|
continue
|
||||||
|
disc_m = _rasterise_polygon(disc_c, image_size, mask_size)
|
||||||
|
if cup_path.exists():
|
||||||
|
cup_c = _load_papila_contour(cup_path)
|
||||||
|
cup_m = (_rasterise_polygon(cup_c, image_size, mask_size)
|
||||||
|
if len(cup_c) >= 3 else np.zeros_like(disc_m))
|
||||||
|
else:
|
||||||
|
cup_m = np.zeros_like(disc_m)
|
||||||
|
discs.append(disc_m)
|
||||||
|
cups.append(cup_m)
|
||||||
|
if not discs:
|
||||||
|
return None
|
||||||
|
disc = (np.mean(discs, axis=0) > 0.5).astype(np.uint8)
|
||||||
|
cup = (np.mean(cups, axis=0) > 0.5).astype(np.uint8)
|
||||||
|
return disc, cup
|
||||||
|
|
||||||
|
|
||||||
|
def _rasterise_polygon(
|
||||||
|
coords: np.ndarray, image_size: Tuple[int, int], target_size: int
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""Rasterise an (N, 2) polygon contour into a (target_size, target_size) binary mask.
|
||||||
|
|
||||||
|
image_size is (width, height) of the coord space (the original fundus image).
|
||||||
|
"""
|
||||||
|
if coords is None or len(coords) < 3:
|
||||||
|
return np.zeros((target_size, target_size), dtype=np.uint8)
|
||||||
|
img = Image.new("L", image_size, 0)
|
||||||
|
pts = [tuple(map(float, p)) for p in coords]
|
||||||
|
ImageDraw.Draw(img).polygon(pts, outline=1, fill=1)
|
||||||
|
img = img.resize((target_size, target_size), Resampling.NEAREST)
|
||||||
|
return (np.array(img, dtype=np.uint8) > 0).astype(np.uint8)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Seg-map loaders
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class GTSegMapLoader:
|
||||||
|
"""Pre-computes per-eye disc/cup seg maps from PAPILA GT contour annotations.
|
||||||
|
|
||||||
|
Output: dict {(pid, eye): np.ndarray (C, H, W) float32} cached for the fold.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
contour_dir: str | Path,
|
||||||
|
*,
|
||||||
|
channels: int = 3,
|
||||||
|
mask_size: int = 512,
|
||||||
|
target_size: int = 224,
|
||||||
|
crop_to_disc: bool = True,
|
||||||
|
) -> None:
|
||||||
|
self._dir = Path(contour_dir)
|
||||||
|
self._channels = channels
|
||||||
|
self._mask_size = mask_size
|
||||||
|
self._target_size = target_size
|
||||||
|
self._crop = crop_to_disc
|
||||||
|
self._cache: dict[tuple, np.ndarray] = {}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cache_dim(self) -> tuple[int, int, int]:
|
||||||
|
return (self._channels, self._target_size, self._target_size)
|
||||||
|
|
||||||
|
def precompute(self, samples: Iterable[Tuple[int, str, Path]]) -> None:
|
||||||
|
"""samples: iterable of (pid, eye, image_path) tuples."""
|
||||||
|
n_ok = 0
|
||||||
|
blank = np.zeros((self._mask_size, self._mask_size), dtype=np.uint8)
|
||||||
|
for pid, eye, image_path in samples:
|
||||||
|
key = (pid, eye)
|
||||||
|
if key in self._cache:
|
||||||
|
continue
|
||||||
|
with Image.open(image_path) as _im:
|
||||||
|
im_size = _im.size # (W, H)
|
||||||
|
res = _papila_disc_cup_masks(
|
||||||
|
pid, eye, self._dir, im_size, self._mask_size,
|
||||||
|
)
|
||||||
|
if res is None:
|
||||||
|
seg = blank
|
||||||
|
else:
|
||||||
|
disc, cup = res
|
||||||
|
seg = _combine_disc_cup(disc, cup)
|
||||||
|
n_ok += 1
|
||||||
|
if self._crop:
|
||||||
|
seg = _crop_to_disc_bbox(seg)
|
||||||
|
self._cache[key] = _seg_map_to_array(seg, self._channels, self._target_size)
|
||||||
|
print(
|
||||||
|
f"[GTSegMapLoader] {n_ok}/{len(self._cache)} GT seg maps computed "
|
||||||
|
f"(channels={self._channels}, target={self._target_size})",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def all_seg_maps(self) -> dict:
|
||||||
|
return self._cache
|
||||||
|
|
||||||
|
|
||||||
|
class _UNetFTDataset(Dataset):
|
||||||
|
"""Pre-cached (image_tensor, mask_tensor) pairs for fine-tuning a UNet.
|
||||||
|
|
||||||
|
Decode + resize + normalize + GT mask rasterisation are all deterministic,
|
||||||
|
so we do them once at construction and store float32 tensors on CPU. This
|
||||||
|
drops per-batch cost to a tensor lookup + GPU transfer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
records: list, # list of (pid, eye, image_path)
|
||||||
|
contour_dir: Path,
|
||||||
|
segmenter, # UNetSegmenter — used for image preprocessing
|
||||||
|
):
|
||||||
|
import time
|
||||||
|
S = segmenter.target_size
|
||||||
|
self._imgs: list[torch.Tensor] = []
|
||||||
|
self._masks: list[torch.Tensor] = []
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"[UNetSegMapLoader] pre-caching {len(records)} (image, mask) pairs "
|
||||||
|
f"at {S}×{S}...",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
t0 = time.time()
|
||||||
|
report = max(1, len(records) // 4)
|
||||||
|
for i, (pid, eye, image_path) in enumerate(records, 1):
|
||||||
|
with Image.open(image_path) as raw:
|
||||||
|
im_size = raw.size
|
||||||
|
img_t = segmenter.preprocess(raw).detach().cpu()
|
||||||
|
res = _papila_disc_cup_masks(pid, eye, contour_dir, im_size, S)
|
||||||
|
if res is None:
|
||||||
|
disc = np.zeros((S, S), dtype=np.uint8)
|
||||||
|
cup = np.zeros_like(disc)
|
||||||
|
else:
|
||||||
|
disc, cup = res
|
||||||
|
mask_t = torch.from_numpy(np.stack([disc, cup], axis=0).astype(np.float32))
|
||||||
|
self._imgs.append(img_t)
|
||||||
|
self._masks.append(mask_t)
|
||||||
|
if i % report == 0 or i == len(records):
|
||||||
|
print(
|
||||||
|
f" [UNet ft cache] {i}/{len(records)} ({time.time() - t0:.1f}s)",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self._imgs)
|
||||||
|
|
||||||
|
def __getitem__(self, idx: int):
|
||||||
|
return self._imgs[idx], self._masks[idx]
|
||||||
|
|
||||||
|
|
||||||
|
class UNetSegMapLoader:
|
||||||
|
"""Pre-computes per-eye seg maps via a REFUGE-pretrained UNet.
|
||||||
|
|
||||||
|
Optionally fine-tunes the UNet per fold on the training split's GT contours.
|
||||||
|
|
||||||
|
Output: dict {(pid, eye): np.ndarray (C, H, W) float32} cached for the fold.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
weights_path: str | Path,
|
||||||
|
*,
|
||||||
|
contour_dir: str | Path,
|
||||||
|
channels: int = 3,
|
||||||
|
target_size: int = 224,
|
||||||
|
unet_size: int = 512,
|
||||||
|
normalize: str = "per_image",
|
||||||
|
threshold: float = 0.5,
|
||||||
|
crop_to_disc: bool = True,
|
||||||
|
finetune_epochs: int = 0,
|
||||||
|
finetune_lr: float = 1e-5,
|
||||||
|
finetune_batch_size: int = 4,
|
||||||
|
device: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
from v4.classes.accessory.unet import UNetSegmenter
|
||||||
|
|
||||||
|
self._weights_path = Path(weights_path)
|
||||||
|
self._contour_dir = Path(contour_dir)
|
||||||
|
self._channels = channels
|
||||||
|
self._target_size = target_size
|
||||||
|
self._threshold = threshold
|
||||||
|
self._crop = crop_to_disc
|
||||||
|
self._ft_epochs = finetune_epochs
|
||||||
|
self._ft_lr = finetune_lr
|
||||||
|
self._ft_batch_size = finetune_batch_size
|
||||||
|
|
||||||
|
self._segmenter = UNetSegmenter(
|
||||||
|
target_size=unet_size, normalize=normalize, device=device,
|
||||||
|
).load_weights(self._weights_path)
|
||||||
|
self._base_state = copy.deepcopy(self._segmenter.model.state_dict())
|
||||||
|
self._cache: dict[tuple, np.ndarray] = {}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cache_dim(self) -> tuple[int, int, int]:
|
||||||
|
return (self._channels, self._target_size, self._target_size)
|
||||||
|
|
||||||
|
def reset_cache(self) -> None:
|
||||||
|
"""Clear cached seg maps (call between folds)."""
|
||||||
|
self._cache.clear()
|
||||||
|
|
||||||
|
def reset_weights(self) -> None:
|
||||||
|
"""Restore base REFUGE weights (undo any prior fine-tuning)."""
|
||||||
|
self._segmenter.model.load_state_dict(copy.deepcopy(self._base_state))
|
||||||
|
|
||||||
|
def finetune(self, train_samples: list) -> None:
|
||||||
|
"""Fine-tune the UNet on the training fold's GT contours.
|
||||||
|
|
||||||
|
train_samples: list of (pid, eye, image_path) tuples — train split only.
|
||||||
|
"""
|
||||||
|
if self._ft_epochs <= 0:
|
||||||
|
return
|
||||||
|
ds = _UNetFTDataset(train_samples, self._contour_dir, self._segmenter)
|
||||||
|
loader = DataLoader(
|
||||||
|
ds, batch_size=self._ft_batch_size, shuffle=True, num_workers=0,
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[UNetSegMapLoader] fine-tuning UNet for {self._ft_epochs} epochs "
|
||||||
|
f"on {len(train_samples)} samples (lr={self._ft_lr}, "
|
||||||
|
f"batch_size={self._ft_batch_size})",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
self._segmenter.finetune(
|
||||||
|
loader, epochs=self._ft_epochs, lr=self._ft_lr,
|
||||||
|
log_prefix="[UNet ft]",
|
||||||
|
)
|
||||||
|
|
||||||
|
def precompute(self, samples: Iterable[Tuple[int, str, Path]]) -> None:
|
||||||
|
"""Run UNet inference on every sample and cache the resulting seg map."""
|
||||||
|
import time
|
||||||
|
samples = list(samples)
|
||||||
|
todo = [s for s in samples if (s[0], s[1]) not in self._cache]
|
||||||
|
if not todo:
|
||||||
|
return
|
||||||
|
print(
|
||||||
|
f"[UNetSegMapLoader] running UNet inference on {len(todo)} images...",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
t0 = time.time()
|
||||||
|
report = max(1, len(todo) // 4)
|
||||||
|
for i, (pid, eye, image_path) in enumerate(todo, 1):
|
||||||
|
with Image.open(image_path) as raw:
|
||||||
|
disc, cup = self._segmenter.predict(raw, threshold=self._threshold)
|
||||||
|
seg = _combine_disc_cup(disc, cup)
|
||||||
|
if self._crop:
|
||||||
|
seg = _crop_to_disc_bbox(seg)
|
||||||
|
self._cache[(pid, eye)] = _seg_map_to_array(
|
||||||
|
seg, self._channels, self._target_size,
|
||||||
|
)
|
||||||
|
if i % report == 0 or i == len(todo):
|
||||||
|
print(
|
||||||
|
f" [UNet inf] {i}/{len(todo)} ({time.time() - t0:.1f}s)",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[UNetSegMapLoader] {len(todo)} seg maps cached via UNet "
|
||||||
|
f"(channels={self._channels}, target={self._target_size})",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def all_seg_maps(self) -> dict:
|
||||||
|
return self._cache
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Factories
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def build_geometry_loader(source: str, **kwargs):
|
||||||
|
"""Return the appropriate geometry-vector loader for the given source string.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
source : "gt" | "unet"
|
||||||
|
contour_dir : (gt) path to contour annotation directory
|
||||||
|
"""
|
||||||
|
if source == "gt":
|
||||||
|
contour_dir = kwargs.get("contour_dir")
|
||||||
|
if contour_dir is None:
|
||||||
|
raise ValueError("build_geometry_loader source='gt' requires contour_dir")
|
||||||
|
return GTGeometryLoader(contour_dir)
|
||||||
|
raise NotImplementedError(f"build_geometry_loader: source={source!r} not implemented")
|
||||||
|
|
||||||
|
|
||||||
|
def build_seg_map_loader(source: str, **kwargs):
|
||||||
|
"""Return the appropriate seg-map loader for the given source string.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
source : "gt" | "unet"
|
||||||
|
|
||||||
|
GT kwargs:
|
||||||
|
contour_dir, channels=3, mask_size=512, target_size=224, crop_to_disc=True
|
||||||
|
UNet kwargs:
|
||||||
|
weights_path, contour_dir, channels=3, target_size=224, unet_size=512,
|
||||||
|
normalize="per_image", threshold=0.5, crop_to_disc=True,
|
||||||
|
finetune_epochs=0, finetune_lr=1e-5, finetune_batch_size=4, device=None
|
||||||
|
"""
|
||||||
|
if source == "gt":
|
||||||
|
if "contour_dir" not in kwargs:
|
||||||
|
raise ValueError("build_seg_map_loader source='gt' requires contour_dir")
|
||||||
|
return GTSegMapLoader(**kwargs)
|
||||||
|
if source == "unet":
|
||||||
|
if "weights_path" not in kwargs:
|
||||||
|
raise ValueError("build_seg_map_loader source='unet' requires weights_path")
|
||||||
|
if "contour_dir" not in kwargs:
|
||||||
|
raise ValueError(
|
||||||
|
"build_seg_map_loader source='unet' requires contour_dir "
|
||||||
|
"(needed for per-fold fine-tuning, even if finetune_epochs=0)"
|
||||||
|
)
|
||||||
|
return UNetSegMapLoader(**kwargs)
|
||||||
|
raise NotImplementedError(f"build_seg_map_loader: source={source!r} not implemented")
|
||||||
@@ -418,6 +418,30 @@ class ImageDataView:
|
|||||||
eye_filter=eye,
|
eye_filter=eye,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ── Geometry hook ─────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _resolve_paths(self, kwargs: dict) -> dict:
|
||||||
|
"""Resolve any *_dir / *_path kwargs against the repo root."""
|
||||||
|
repo_root = Path(__file__).resolve().parents[3]
|
||||||
|
out = {}
|
||||||
|
for k, v in kwargs.items():
|
||||||
|
if (k.endswith("_dir") or k.endswith("_path")) and v is not None:
|
||||||
|
p = Path(v)
|
||||||
|
out[k] = str(repo_root / p) if not p.is_absolute() else v
|
||||||
|
else:
|
||||||
|
out[k] = v
|
||||||
|
return out
|
||||||
|
|
||||||
|
def build_geometry_loader(self, source: str, **kwargs):
|
||||||
|
"""Return a geometry-vector loader (delegates to fundus_images)."""
|
||||||
|
from v4.classes.profiles.fundus_images import build_geometry_loader as _build
|
||||||
|
return _build(source, **self._resolve_paths(kwargs))
|
||||||
|
|
||||||
|
def build_seg_map_loader(self, source: str, **kwargs):
|
||||||
|
"""Return a seg-map loader (delegates to fundus_images)."""
|
||||||
|
from v4.classes.profiles.fundus_images import build_seg_map_loader as _build
|
||||||
|
return _build(source, **self._resolve_paths(kwargs))
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# PapilaBundle — the v4 DataBundle returned by build_data
|
# PapilaBundle — the v4 DataBundle returned by build_data
|
||||||
@@ -515,6 +539,7 @@ class PapilaBundle:
|
|||||||
*,
|
*,
|
||||||
level: str = "eye",
|
level: str = "eye",
|
||||||
label_filter: list[int] | None = None,
|
label_filter: list[int] | None = None,
|
||||||
|
eye_filter: str | None = None,
|
||||||
) -> LoaderShell:
|
) -> LoaderShell:
|
||||||
"""Build a LoaderShell from a split DataFrame.
|
"""Build a LoaderShell from a split DataFrame.
|
||||||
|
|
||||||
@@ -533,6 +558,8 @@ class PapilaBundle:
|
|||||||
|
|
||||||
if label_filter is not None:
|
if label_filter is not None:
|
||||||
df = df[df[lc].isin(label_filter)]
|
df = df[df[lc].isin(label_filter)]
|
||||||
|
if eye_filter is not None and "eyeID" in df.columns:
|
||||||
|
df = df[df["eyeID"] == eye_filter]
|
||||||
|
|
||||||
entries: list[ShellEntry] = []
|
entries: list[ShellEntry] = []
|
||||||
|
|
||||||
|
|||||||
+60
-17
@@ -11,7 +11,8 @@ import torch.nn.functional as F
|
|||||||
from v4.classes.dataset import LoaderShell, to_label_tensor
|
from v4.classes.dataset import LoaderShell, to_label_tensor
|
||||||
from v4.classes.metrics import score_arrays, compute_extended_metrics, tune_binary_threshold
|
from v4.classes.metrics import score_arrays, compute_extended_metrics, tune_binary_threshold
|
||||||
from v4.classes.stages.helpers import (
|
from v4.classes.stages.helpers import (
|
||||||
encode_embedding, get_out_dim, resolve_input_dims, phase_for_epoch,
|
class_weights_from_shell, encode_embedding, get_out_dim, resolve_input_dims,
|
||||||
|
phase_for_epoch,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -25,15 +26,15 @@ def collect_probs(
|
|||||||
loader,
|
loader,
|
||||||
device,
|
device,
|
||||||
num_classes: int,
|
num_classes: int,
|
||||||
) -> tuple[np.ndarray, np.ndarray]:
|
) -> tuple[np.ndarray, np.ndarray, list, np.ndarray]:
|
||||||
"""Eval pass for one fusion stage; returns (y_true, softmax_probs)."""
|
"""Eval pass for one fusion stage; returns (y_true, softmax_probs, entity_ids, embeddings)."""
|
||||||
from v4.classes.dataset import to_label_tensor
|
from v4.classes.dataset import to_label_tensor
|
||||||
bridge.eval(); primary_head.eval()
|
bridge.eval(); primary_head.eval()
|
||||||
for t in towers.values():
|
for t in towers.values():
|
||||||
t.eval()
|
t.eval()
|
||||||
inputs = stage_cfg["inputs"]
|
inputs = stage_cfg["inputs"]
|
||||||
is_bilateral = isinstance(inputs, dict)
|
is_bilateral = isinstance(inputs, dict)
|
||||||
y_all, p_all = [], []
|
y_all, p_all, ids_all, z_all = [], [], [], []
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for batch in loader:
|
for batch in loader:
|
||||||
@@ -57,10 +58,14 @@ def collect_probs(
|
|||||||
logits = primary_head(z)
|
logits = primary_head(z)
|
||||||
y_all.append(to_label_tensor(y, device).cpu().numpy())
|
y_all.append(to_label_tensor(y, device).cpu().numpy())
|
||||||
p_all.append(F.softmax(logits, dim=1).cpu().numpy())
|
p_all.append(F.softmax(logits, dim=1).cpu().numpy())
|
||||||
|
z_all.append(z.cpu().numpy())
|
||||||
|
ids_all.extend(batch.get("entity_id", []))
|
||||||
|
|
||||||
if not y_all:
|
if not y_all:
|
||||||
return np.zeros(0, dtype=np.int64), np.zeros((0, num_classes), dtype=np.float32)
|
return (np.zeros(0, dtype=np.int64), np.zeros((0, num_classes), dtype=np.float32),
|
||||||
return np.concatenate(y_all), np.concatenate(p_all, axis=0)
|
[], np.zeros((0, 0), dtype=np.float32))
|
||||||
|
return (np.concatenate(y_all), np.concatenate(p_all, axis=0),
|
||||||
|
ids_all, np.concatenate(z_all, axis=0))
|
||||||
|
|
||||||
|
|
||||||
def run(
|
def run(
|
||||||
@@ -88,9 +93,13 @@ def run(
|
|||||||
inputs = stage_cfg["inputs"]
|
inputs = stage_cfg["inputs"]
|
||||||
is_bilateral = isinstance(inputs, dict)
|
is_bilateral = isinstance(inputs, dict)
|
||||||
|
|
||||||
s_train = data.build_shells(split.train, level=level, label_filter=label_filter)
|
eye_filter = stage_cfg.get("eye_filter", None)
|
||||||
s_val = data.build_shells(split.val, level=level, label_filter=label_filter)
|
s_train = data.build_shells(split.train, level=level, label_filter=label_filter,
|
||||||
s_test = (data.build_shells(split.test, level=level, label_filter=label_filter)
|
eye_filter=eye_filter)
|
||||||
|
s_val = data.build_shells(split.val, level=level, label_filter=label_filter,
|
||||||
|
eye_filter=eye_filter)
|
||||||
|
s_test = (data.build_shells(split.test, level=level, label_filter=label_filter,
|
||||||
|
eye_filter=eye_filter)
|
||||||
if split.test is not None else LoaderShell(entries=[]))
|
if split.test is not None else LoaderShell(entries=[]))
|
||||||
|
|
||||||
if not s_val.entries:
|
if not s_val.entries:
|
||||||
@@ -116,7 +125,11 @@ def run(
|
|||||||
h_dim = get_out_dim(hs["input"], towers, {**stage_models, name: bridge})
|
h_dim = get_out_dim(hs["input"], towers, {**stage_models, name: bridge})
|
||||||
h_mod = importlib.import_module(hs.get("module", "v4.classes.heads.classifier"))
|
h_mod = importlib.import_module(hs.get("module", "v4.classes.heads.classifier"))
|
||||||
h_cls = getattr(h_mod, hs.get("class", "ClassificationHead"))
|
h_cls = getattr(h_mod, hs.get("class", "ClassificationHead"))
|
||||||
head_models[hs["name"]] = h_cls(h_dim, num_classes).to(device)
|
existing = stage_models.get(hs["name"])
|
||||||
|
head_models[hs["name"]] = (
|
||||||
|
existing.to(device) if existing is not None
|
||||||
|
else h_cls(h_dim, num_classes, **hs.get("args", {})).to(device)
|
||||||
|
)
|
||||||
|
|
||||||
primary_hs_cfg = next((hs for hs in head_stage_cfgs if not hs.get("bcd", False)), None)
|
primary_hs_cfg = next((hs for hs in head_stage_cfgs if not hs.get("bcd", False)), None)
|
||||||
bcd_head_cfgs = [hs for hs in head_stage_cfgs if hs.get("bcd", False)]
|
bcd_head_cfgs = [hs for hs in head_stage_cfgs if hs.get("bcd", False)]
|
||||||
@@ -132,11 +145,23 @@ def run(
|
|||||||
for p in m.parameters():
|
for p in m.parameters():
|
||||||
p.requires_grad_(False)
|
p.requires_grad_(False)
|
||||||
m.eval()
|
m.eval()
|
||||||
|
for h in head_models.values():
|
||||||
|
for p in h.parameters():
|
||||||
|
p.requires_grad_(True)
|
||||||
|
|
||||||
# ── Optimizer ────────────────────────────────────────────────────────────
|
# ── Optimizer ────────────────────────────────────────────────────────────
|
||||||
train_towers = stage_cfg.get("train_towers", False)
|
train_towers = stage_cfg.get("train_towers", False)
|
||||||
|
if train_towers:
|
||||||
|
# Only train towers that are direct inputs to this stage (not all towers globally).
|
||||||
|
# For nt_od with inputs ["img_od", "cd_od"] this trains only those two; other
|
||||||
|
# eye's towers remain untouched.
|
||||||
|
direct_inputs = list(inputs.values()) if isinstance(inputs, dict) else inputs
|
||||||
|
tower_params = [p for n in direct_inputs if n in towers
|
||||||
|
for p in towers[n].parameters()]
|
||||||
|
else:
|
||||||
|
tower_params = []
|
||||||
opt_params = (
|
opt_params = (
|
||||||
([p for t in towers.values() for p in t.parameters()] if train_towers else []) +
|
tower_params +
|
||||||
list(bridge.parameters()) +
|
list(bridge.parameters()) +
|
||||||
[p for h in head_models.values() for p in h.parameters()]
|
[p for h in head_models.values() for p in h.parameters()]
|
||||||
)
|
)
|
||||||
@@ -147,6 +172,17 @@ def run(
|
|||||||
wf = 0 if is_bilateral else warmup_cfg.get("fused_epochs", 0)
|
wf = 0 if is_bilateral else warmup_cfg.get("fused_epochs", 0)
|
||||||
bcd_prob = cfg["training"].get("bcd_prob", 0.5)
|
bcd_prob = cfg["training"].get("bcd_prob", 0.5)
|
||||||
|
|
||||||
|
cw = class_weights_from_shell(
|
||||||
|
s_train, num_classes, device,
|
||||||
|
enabled=cfg["training"].get("class_weighted", False),
|
||||||
|
)
|
||||||
|
if cw is not None:
|
||||||
|
print(
|
||||||
|
f" fold{fold+1} [{name}] class weights: "
|
||||||
|
+ ", ".join(f"{i}={w:.3f}" for i, w in enumerate(cw.tolist())),
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
# ── Epoch loop ────────────────────────────────────────────────────────────
|
# ── Epoch loop ────────────────────────────────────────────────────────────
|
||||||
for epoch in range(epochs):
|
for epoch in range(epochs):
|
||||||
bridge.train()
|
bridge.train()
|
||||||
@@ -200,7 +236,7 @@ def run(
|
|||||||
if is_bilateral or phase == "fused_warmup":
|
if is_bilateral or phase == "fused_warmup":
|
||||||
logits = head_logits.get(primary_hs_cfg["name"])
|
logits = head_logits.get(primary_hs_cfg["name"])
|
||||||
elif phase == "tower_warmup" and bcd_head_cfgs:
|
elif phase == "tower_warmup" and bcd_head_cfgs:
|
||||||
losses = [F.cross_entropy(head_logits[hs["name"]], y_t)
|
losses = [F.cross_entropy(head_logits[hs["name"]], y_t, weight=cw)
|
||||||
for hs in bcd_head_cfgs if hs["name"] in head_logits]
|
for hs in bcd_head_cfgs if hs["name"] in head_logits]
|
||||||
if not losses:
|
if not losses:
|
||||||
continue
|
continue
|
||||||
@@ -217,7 +253,7 @@ def run(
|
|||||||
|
|
||||||
if logits is None:
|
if logits is None:
|
||||||
continue
|
continue
|
||||||
loss = F.cross_entropy(logits, y_t)
|
loss = F.cross_entropy(logits, y_t, weight=cw)
|
||||||
opt.zero_grad(); loss.backward(); opt.step()
|
opt.zero_grad(); loss.backward(); opt.step()
|
||||||
total_correct += int((logits.argmax(1) == y_t).sum())
|
total_correct += int((logits.argmax(1) == y_t).sum())
|
||||||
total_loss += loss.item() * len(y_t)
|
total_loss += loss.item() * len(y_t)
|
||||||
@@ -226,7 +262,7 @@ def run(
|
|||||||
tr_loss = total_loss / total_n if total_n else nan
|
tr_loss = total_loss / total_n if total_n else nan
|
||||||
tr_acc = total_correct / total_n if total_n else nan
|
tr_acc = total_correct / total_n if total_n else nan
|
||||||
|
|
||||||
y_v, p_v = collect_probs(bridge, primary_head, stage_cfg, towers,
|
y_v, p_v, _, _ = collect_probs(bridge, primary_head, stage_cfg, towers,
|
||||||
stage_models, cfg_stages, val_loader, device, num_classes)
|
stage_models, cfg_stages, val_loader, device, num_classes)
|
||||||
_, val_auc, _ = score_arrays(y_v, p_v, num_classes) if y_v.size else (nan, nan, nan)
|
_, val_auc, _ = score_arrays(y_v, p_v, num_classes) if y_v.size else (nan, nan, nan)
|
||||||
print(
|
print(
|
||||||
@@ -236,7 +272,7 @@ def run(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# ── Final eval ────────────────────────────────────────────────────────────
|
# ── Final eval ────────────────────────────────────────────────────────────
|
||||||
y_val, p_val = collect_probs(bridge, primary_head, stage_cfg, towers,
|
y_val, p_val, ids_val, z_val = collect_probs(bridge, primary_head, stage_cfg, towers,
|
||||||
stage_models, cfg_stages, val_loader, device, num_classes)
|
stage_models, cfg_stages, val_loader, device, num_classes)
|
||||||
val_acc, val_auc, val_n = (score_arrays(y_val, p_val, num_classes)
|
val_acc, val_auc, val_n = (score_arrays(y_val, p_val, num_classes)
|
||||||
if y_val.size else (nan, nan, nan))
|
if y_val.size else (nan, nan, nan))
|
||||||
@@ -247,9 +283,10 @@ def run(
|
|||||||
and num_classes == 2 and y_val.size >= 2):
|
and num_classes == 2 and y_val.size >= 2):
|
||||||
val_threshold = tune_binary_threshold(y_val, p_val[:, 1])
|
val_threshold = tune_binary_threshold(y_val, p_val[:, 1])
|
||||||
|
|
||||||
|
y_te = p_te = ids_te = z_te = None
|
||||||
test_auc = test_acc = test_n = nan
|
test_auc = test_acc = test_n = nan
|
||||||
if test_loader is not None:
|
if test_loader is not None:
|
||||||
y_te, p_te = collect_probs(bridge, primary_head, stage_cfg, towers,
|
y_te, p_te, ids_te, z_te = collect_probs(bridge, primary_head, stage_cfg, towers,
|
||||||
stage_models, cfg_stages, test_loader, device, num_classes)
|
stage_models, cfg_stages, test_loader, device, num_classes)
|
||||||
test_acc, test_auc, test_n = (score_arrays(y_te, p_te, num_classes)
|
test_acc, test_auc, test_n = (score_arrays(y_te, p_te, num_classes)
|
||||||
if y_te.size else (nan, nan, nan))
|
if y_te.size else (nan, nan, nan))
|
||||||
@@ -270,4 +307,10 @@ def run(
|
|||||||
f"{name}_test_acc": test_acc,
|
f"{name}_test_acc": test_acc,
|
||||||
f"{name}_test_n": test_n,
|
f"{name}_test_n": test_n,
|
||||||
}
|
}
|
||||||
return updated, metrics
|
pred_data = {
|
||||||
|
name: {
|
||||||
|
"val_y": y_val, "val_p": p_val, "val_ids": ids_val, "val_z": z_val,
|
||||||
|
"test_y": y_te, "test_p": p_te, "test_ids": ids_te, "test_z": z_te,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return updated, metrics, pred_data
|
||||||
|
|||||||
@@ -1,9 +1,32 @@
|
|||||||
"""stages/helpers — shared utilities for stage runners."""
|
"""stages/helpers — shared utilities for stage runners."""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections import Counter
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
def class_weights_from_shell(
|
||||||
|
shell, num_classes: int, device, *, enabled: bool = True
|
||||||
|
) -> torch.Tensor | None:
|
||||||
|
"""Return inverse-frequency CE weights normalised to mean 1, or None.
|
||||||
|
|
||||||
|
weights[i] = (N_total / num_classes) / N_class_i → rare classes weighted higher.
|
||||||
|
Mean(weights) ≈ 1 so overall loss magnitude is unchanged.
|
||||||
|
|
||||||
|
Classes absent from the shell get weight 1.0 (no division-by-zero).
|
||||||
|
"""
|
||||||
|
if not enabled or shell is None or not shell.entries:
|
||||||
|
return None
|
||||||
|
counts = Counter(int(e.label) for e in shell.entries)
|
||||||
|
n_total = sum(counts.values())
|
||||||
|
weights = []
|
||||||
|
for c in range(num_classes):
|
||||||
|
n_c = counts.get(c, 0)
|
||||||
|
weights.append(1.0 if n_c == 0 else n_total / (num_classes * n_c))
|
||||||
|
return torch.tensor(weights, dtype=torch.float32, device=device)
|
||||||
|
|
||||||
|
|
||||||
def get_out_dim(name: str, towers: dict, stage_models: dict) -> int:
|
def get_out_dim(name: str, towers: dict, stage_models: dict) -> int:
|
||||||
if name in towers:
|
if name in towers:
|
||||||
return towers[name].out_dim
|
return towers[name].out_dim
|
||||||
|
|||||||
@@ -0,0 +1,394 @@
|
|||||||
|
"""stages/parallel — runs multiple same-type sub-stages in a shared epoch loop.
|
||||||
|
|
||||||
|
Used to train bilateral pairs (OD + OS) simultaneously rather than sequentially.
|
||||||
|
Sub-stages must all be the same type: either all 'warm' or all 'fusion'.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import importlib
|
||||||
|
from random import choice, random as _random
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from v4.classes.dataset import LoaderShell, to_label_tensor
|
||||||
|
from v4.classes.metrics import score_arrays, compute_extended_metrics, tune_binary_threshold
|
||||||
|
from v4.classes.stages.helpers import (
|
||||||
|
class_weights_from_shell, get_out_dim, resolve_input_dims, phase_for_epoch,
|
||||||
|
)
|
||||||
|
from v4.classes.stages.fusion import collect_probs
|
||||||
|
|
||||||
|
|
||||||
|
def run(
|
||||||
|
stage_cfg: dict,
|
||||||
|
cfg: dict,
|
||||||
|
towers: dict,
|
||||||
|
stage_models: dict,
|
||||||
|
data,
|
||||||
|
split,
|
||||||
|
label_filter,
|
||||||
|
num_classes: int,
|
||||||
|
device,
|
||||||
|
fold: int,
|
||||||
|
cfg_stages: list[dict],
|
||||||
|
_make_loader,
|
||||||
|
_balanced_sampler,
|
||||||
|
) -> tuple[dict, dict, dict]:
|
||||||
|
sub_cfgs = stage_cfg["stages"]
|
||||||
|
sub_types = {s["type"] for s in sub_cfgs}
|
||||||
|
|
||||||
|
if sub_types == {"warm"}:
|
||||||
|
return _parallel_warm(sub_cfgs, cfg, towers, stage_models, data, split,
|
||||||
|
label_filter, num_classes, device, fold,
|
||||||
|
cfg_stages, _make_loader, _balanced_sampler)
|
||||||
|
elif sub_types == {"fusion"}:
|
||||||
|
return _parallel_fusion(sub_cfgs, cfg, towers, stage_models, data, split,
|
||||||
|
label_filter, num_classes, device, fold,
|
||||||
|
cfg_stages, _make_loader)
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"parallel stage sub-stages must all be the same type (warm or fusion), "
|
||||||
|
f"got {sub_types}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── Parallel warm ─────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _parallel_warm(
|
||||||
|
sub_cfgs, cfg, towers, stage_models, data, split,
|
||||||
|
label_filter, num_classes, device, fold, cfg_stages, _make_loader, _balanced_sampler,
|
||||||
|
):
|
||||||
|
bs = cfg["training"]["batch_size"]
|
||||||
|
|
||||||
|
contexts = []
|
||||||
|
for sc in sub_cfgs:
|
||||||
|
tower_name = sc["tower"]
|
||||||
|
level = sc["level"]
|
||||||
|
shell_filter = sc.get("shell_filter", {})
|
||||||
|
n_epochs = sc.get("epochs", 0)
|
||||||
|
|
||||||
|
s_train = data.build_shells(split.train, level=level, label_filter=label_filter,
|
||||||
|
**shell_filter)
|
||||||
|
loader = _make_loader(s_train, {tower_name: towers[tower_name]},
|
||||||
|
batch_size=bs, shuffle=False,
|
||||||
|
sampler=_balanced_sampler(s_train))
|
||||||
|
head_name = sc.get("head_name")
|
||||||
|
if head_name:
|
||||||
|
head_cfg = next((s for s in cfg_stages if s.get("name") == head_name), None)
|
||||||
|
if head_cfg is None:
|
||||||
|
raise ValueError(f"warm stage requested head_name={head_name!r}, but no such head exists")
|
||||||
|
h_mod = importlib.import_module(head_cfg.get("module", "v4.classes.heads.classifier"))
|
||||||
|
h_cls = getattr(h_mod, head_cfg.get("class", "ClassificationHead"))
|
||||||
|
probe = stage_models.get(head_name)
|
||||||
|
if probe is None:
|
||||||
|
probe = h_cls(towers[tower_name].out_dim, num_classes, **head_cfg.get("args", {}))
|
||||||
|
probe = probe.to(device)
|
||||||
|
else:
|
||||||
|
probe = torch.nn.Linear(towers[tower_name].out_dim, num_classes).to(device)
|
||||||
|
opt = torch.optim.Adam(
|
||||||
|
list(towers[tower_name].parameters()) + list(probe.parameters()),
|
||||||
|
lr=cfg["training"]["lr"],
|
||||||
|
)
|
||||||
|
cw = class_weights_from_shell(
|
||||||
|
s_train, num_classes, device,
|
||||||
|
enabled=cfg["training"].get("class_weighted", False),
|
||||||
|
)
|
||||||
|
contexts.append({
|
||||||
|
"name": tower_name, "n_epochs": n_epochs,
|
||||||
|
"loader": loader, "probe": probe, "opt": opt,
|
||||||
|
"head_name": head_name, "class_weights": cw,
|
||||||
|
})
|
||||||
|
|
||||||
|
active = {c["name"] for c in contexts if c["n_epochs"] > 0}
|
||||||
|
for name, t in towers.items():
|
||||||
|
for p in t.parameters():
|
||||||
|
p.requires_grad_(name in active)
|
||||||
|
|
||||||
|
max_epochs = max((c["n_epochs"] for c in contexts), default=0)
|
||||||
|
for epoch in range(max_epochs):
|
||||||
|
for ctx in contexts:
|
||||||
|
if epoch >= ctx["n_epochs"]:
|
||||||
|
continue
|
||||||
|
tower_name = ctx["name"]
|
||||||
|
towers[tower_name].train()
|
||||||
|
total_loss = total_correct = total_n = 0
|
||||||
|
for batch in ctx["loader"]:
|
||||||
|
y = batch.get("label")
|
||||||
|
x = batch.get(tower_name)
|
||||||
|
if not torch.is_tensor(y) or not torch.is_tensor(x):
|
||||||
|
continue
|
||||||
|
y_t = to_label_tensor(y, device)
|
||||||
|
logits = ctx["probe"](towers[tower_name](x.to(device)))
|
||||||
|
loss = F.cross_entropy(logits, y_t, weight=ctx["class_weights"])
|
||||||
|
ctx["opt"].zero_grad(); loss.backward(); ctx["opt"].step()
|
||||||
|
total_loss += loss.item() * len(y_t)
|
||||||
|
total_correct += int((logits.argmax(1) == y_t).sum())
|
||||||
|
total_n += len(y_t)
|
||||||
|
if total_n:
|
||||||
|
print(
|
||||||
|
f" fold{fold+1} [warm/{tower_name}]"
|
||||||
|
f" ep{epoch+1:03d}/{ctx['n_epochs']}"
|
||||||
|
f" loss={total_loss/total_n:.4f} acc={total_correct/total_n:.3f}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
for t in towers.values():
|
||||||
|
for p in t.parameters():
|
||||||
|
p.requires_grad_(True)
|
||||||
|
|
||||||
|
updated = dict(stage_models)
|
||||||
|
for ctx in contexts:
|
||||||
|
if ctx.get("head_name"):
|
||||||
|
updated[ctx["head_name"]] = ctx["probe"]
|
||||||
|
|
||||||
|
return updated, {}, {}
|
||||||
|
|
||||||
|
|
||||||
|
# ── Parallel fusion ───────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _parallel_fusion(
|
||||||
|
sub_cfgs, cfg, towers, stage_models, data, split,
|
||||||
|
label_filter, num_classes, device, fold, cfg_stages, _make_loader,
|
||||||
|
):
|
||||||
|
nan = float("nan")
|
||||||
|
bs = cfg["training"]["batch_size"]
|
||||||
|
bcd_prob = cfg["training"].get("bcd_prob", 0.5)
|
||||||
|
|
||||||
|
# Freeze all prior stage models once, before building any bridges.
|
||||||
|
for m in stage_models.values():
|
||||||
|
for p in m.parameters():
|
||||||
|
p.requires_grad_(False)
|
||||||
|
m.eval()
|
||||||
|
|
||||||
|
# ── Per-sub-stage setup ───────────────────────────────────────────────────
|
||||||
|
contexts = []
|
||||||
|
for sc in sub_cfgs:
|
||||||
|
name = sc["name"]
|
||||||
|
level = sc["level"]
|
||||||
|
inputs = sc["inputs"]
|
||||||
|
epochs = sc["epochs"]
|
||||||
|
eye_filter = sc.get("eye_filter", None)
|
||||||
|
|
||||||
|
s_train = data.build_shells(split.train, level=level, label_filter=label_filter,
|
||||||
|
eye_filter=eye_filter)
|
||||||
|
s_val = data.build_shells(split.val, level=level, label_filter=label_filter,
|
||||||
|
eye_filter=eye_filter)
|
||||||
|
s_test = (data.build_shells(split.test, level=level, label_filter=label_filter,
|
||||||
|
eye_filter=eye_filter)
|
||||||
|
if split.test is not None else LoaderShell(entries=[]))
|
||||||
|
|
||||||
|
if not s_val.entries:
|
||||||
|
print(f" fold{fold+1}: no val samples for stage {name!r}, skipping.", flush=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
input_dims = resolve_input_dims(inputs, towers, stage_models)
|
||||||
|
bmod = importlib.import_module(sc["module"])
|
||||||
|
bridge = getattr(bmod, sc["class"])(input_dims, **sc.get("args", {})).to(device)
|
||||||
|
|
||||||
|
# Head stages for this sub-stage.
|
||||||
|
head_stage_cfgs = [s for s in cfg_stages if s["type"] == "head"
|
||||||
|
and s.get("train_with") == name]
|
||||||
|
head_models: dict[str, torch.nn.Module] = {}
|
||||||
|
for hs in head_stage_cfgs:
|
||||||
|
h_dim = get_out_dim(hs["input"], towers, {**stage_models, name: bridge})
|
||||||
|
h_mod = importlib.import_module(hs.get("module", "v4.classes.heads.classifier"))
|
||||||
|
h_cls = getattr(h_mod, hs.get("class", "ClassificationHead"))
|
||||||
|
existing = stage_models.get(hs["name"])
|
||||||
|
head_models[hs["name"]] = (
|
||||||
|
existing.to(device) if existing is not None
|
||||||
|
else h_cls(h_dim, num_classes, **hs.get("args", {})).to(device)
|
||||||
|
)
|
||||||
|
for h in head_models.values():
|
||||||
|
for p in h.parameters():
|
||||||
|
p.requires_grad_(True)
|
||||||
|
|
||||||
|
primary_hs_cfg = next((hs for hs in head_stage_cfgs if not hs.get("bcd", False)), None)
|
||||||
|
bcd_head_cfgs = [hs for hs in head_stage_cfgs if hs.get("bcd", False)]
|
||||||
|
|
||||||
|
if primary_hs_cfg is None:
|
||||||
|
print(f" WARNING: no primary head for stage {name!r}; skipping.", flush=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
primary_head = head_models[primary_hs_cfg["name"]]
|
||||||
|
|
||||||
|
train_towers = sc.get("train_towers", False)
|
||||||
|
direct_inputs = inputs if isinstance(inputs, list) else list(inputs.values())
|
||||||
|
tower_params = ([p for n in direct_inputs if n in towers
|
||||||
|
for p in towers[n].parameters()]
|
||||||
|
if train_towers else [])
|
||||||
|
opt_params = (tower_params + list(bridge.parameters()) +
|
||||||
|
[p for h in head_models.values() for p in h.parameters()])
|
||||||
|
opt = torch.optim.Adam(opt_params, lr=cfg["training"]["lr"])
|
||||||
|
|
||||||
|
warmup_cfg = sc.get("warmup", {})
|
||||||
|
wt = warmup_cfg.get("tower_epochs", 0)
|
||||||
|
wf = warmup_cfg.get("fused_epochs", 0)
|
||||||
|
|
||||||
|
train_loader = _make_loader(s_train, towers, batch_size=bs, shuffle=True)
|
||||||
|
val_loader = _make_loader(s_val, towers, batch_size=bs, shuffle=False)
|
||||||
|
test_loader = (_make_loader(s_test, towers, batch_size=bs, shuffle=False)
|
||||||
|
if s_test.entries else None)
|
||||||
|
|
||||||
|
cw = class_weights_from_shell(
|
||||||
|
s_train, num_classes, device,
|
||||||
|
enabled=cfg["training"].get("class_weighted", False),
|
||||||
|
)
|
||||||
|
contexts.append({
|
||||||
|
"name": name, "inputs": inputs, "epochs": epochs,
|
||||||
|
"bridge": bridge, "head_models": head_models,
|
||||||
|
"primary_head": primary_head, "primary_hs_cfg": primary_hs_cfg,
|
||||||
|
"bcd_head_cfgs": bcd_head_cfgs,
|
||||||
|
"opt": opt, "wt": wt, "wf": wf,
|
||||||
|
"train_towers": train_towers, "direct_inputs": direct_inputs,
|
||||||
|
"train_loader": train_loader, "val_loader": val_loader,
|
||||||
|
"test_loader": test_loader, "sc": sc,
|
||||||
|
"class_weights": cw,
|
||||||
|
})
|
||||||
|
|
||||||
|
if not contexts:
|
||||||
|
return stage_models, {}, {}
|
||||||
|
|
||||||
|
max_epochs = max(c["epochs"] for c in contexts)
|
||||||
|
|
||||||
|
# ── Shared epoch loop ─────────────────────────────────────────────────────
|
||||||
|
for epoch in range(max_epochs):
|
||||||
|
for ctx in contexts:
|
||||||
|
if epoch >= ctx["epochs"]:
|
||||||
|
continue
|
||||||
|
|
||||||
|
name = ctx["name"]
|
||||||
|
bridge = ctx["bridge"]
|
||||||
|
inputs = ctx["inputs"]
|
||||||
|
wt, wf = ctx["wt"], ctx["wf"]
|
||||||
|
phase = phase_for_epoch(epoch, wt, wf)
|
||||||
|
|
||||||
|
bridge.train()
|
||||||
|
for h in ctx["head_models"].values():
|
||||||
|
h.train()
|
||||||
|
if ctx["train_towers"]:
|
||||||
|
for n in ctx["direct_inputs"]:
|
||||||
|
if n in towers:
|
||||||
|
towers[n].train()
|
||||||
|
else:
|
||||||
|
for n in ctx["direct_inputs"]:
|
||||||
|
if n in towers:
|
||||||
|
towers[n].eval()
|
||||||
|
|
||||||
|
if hasattr(bridge, "set_phase"):
|
||||||
|
bridge.set_phase(phase)
|
||||||
|
|
||||||
|
total_loss = total_correct = total_n = 0
|
||||||
|
|
||||||
|
for batch in ctx["train_loader"]:
|
||||||
|
y = batch.get("label")
|
||||||
|
if not torch.is_tensor(y):
|
||||||
|
continue
|
||||||
|
y_t = to_label_tensor(y, device)
|
||||||
|
if y_t.numel() == 0:
|
||||||
|
continue
|
||||||
|
|
||||||
|
local_embs = {n: towers[n](batch[n].to(device)) for n in inputs
|
||||||
|
if n in batch and torch.is_tensor(batch[n])}
|
||||||
|
if len(local_embs) != len(inputs):
|
||||||
|
continue
|
||||||
|
local_embs[name] = bridge(list(local_embs[n] for n in inputs))
|
||||||
|
|
||||||
|
head_logits = {
|
||||||
|
hs["name"]: ctx["head_models"][hs["name"]](local_embs[hs["input"]])
|
||||||
|
for hs in ([ctx["primary_hs_cfg"]] + ctx["bcd_head_cfgs"])
|
||||||
|
if hs["input"] in local_embs
|
||||||
|
}
|
||||||
|
|
||||||
|
if phase == "fused_warmup":
|
||||||
|
logits = head_logits.get(ctx["primary_hs_cfg"]["name"])
|
||||||
|
elif phase == "tower_warmup" and ctx["bcd_head_cfgs"]:
|
||||||
|
losses = [F.cross_entropy(head_logits[hs["name"]], y_t,
|
||||||
|
weight=ctx["class_weights"])
|
||||||
|
for hs in ctx["bcd_head_cfgs"] if hs["name"] in head_logits]
|
||||||
|
if not losses:
|
||||||
|
continue
|
||||||
|
loss = sum(losses) / len(losses)
|
||||||
|
ctx["opt"].zero_grad(); loss.backward(); ctx["opt"].step()
|
||||||
|
total_loss += loss.item() * len(y_t)
|
||||||
|
total_n += len(y_t)
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
if ctx["bcd_head_cfgs"] and _random() < bcd_prob:
|
||||||
|
logits = head_logits.get(choice(ctx["bcd_head_cfgs"])["name"])
|
||||||
|
else:
|
||||||
|
logits = head_logits.get(ctx["primary_hs_cfg"]["name"])
|
||||||
|
|
||||||
|
if logits is None:
|
||||||
|
continue
|
||||||
|
loss = F.cross_entropy(logits, y_t, weight=ctx["class_weights"])
|
||||||
|
ctx["opt"].zero_grad(); loss.backward(); ctx["opt"].step()
|
||||||
|
total_correct += int((logits.argmax(1) == y_t).sum())
|
||||||
|
total_loss += loss.item() * len(y_t)
|
||||||
|
total_n += len(y_t)
|
||||||
|
|
||||||
|
tr_loss = total_loss / total_n if total_n else nan
|
||||||
|
tr_acc = total_correct / total_n if total_n else nan
|
||||||
|
|
||||||
|
y_v, p_v, _, _ = collect_probs(bridge, ctx["primary_head"], ctx["sc"], towers,
|
||||||
|
stage_models, cfg_stages, ctx["val_loader"],
|
||||||
|
device, num_classes)
|
||||||
|
_, val_auc, _ = score_arrays(y_v, p_v, num_classes) if y_v.size else (nan, nan, nan)
|
||||||
|
print(
|
||||||
|
f" fold{fold+1} [{name}] ep{epoch+1:03d}/{ctx['epochs']} [{phase:14s}]"
|
||||||
|
f" loss={tr_loss:.4f} acc={tr_acc:.3f} val_auc={val_auc:.4f}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Final eval + collect results ──────────────────────────────────────────
|
||||||
|
updated = dict(stage_models)
|
||||||
|
all_metrics: dict = {}
|
||||||
|
all_preds: dict = {}
|
||||||
|
|
||||||
|
for ctx in contexts:
|
||||||
|
name = ctx["name"]
|
||||||
|
bridge = ctx["bridge"]
|
||||||
|
primary_head = ctx["primary_head"]
|
||||||
|
|
||||||
|
updated[name] = bridge
|
||||||
|
updated.update(ctx["head_models"])
|
||||||
|
|
||||||
|
y_val, p_val, ids_val, z_val = collect_probs(bridge, primary_head, ctx["sc"], towers,
|
||||||
|
stage_models, cfg_stages, ctx["val_loader"],
|
||||||
|
device, num_classes)
|
||||||
|
val_acc, val_auc, val_n = (score_arrays(y_val, p_val, num_classes)
|
||||||
|
if y_val.size else (nan, nan, nan))
|
||||||
|
ext = compute_extended_metrics(y_val, p_val, num_classes) if y_val.size else {}
|
||||||
|
|
||||||
|
val_threshold = 0.5
|
||||||
|
if (cfg["training"].get("tune_binary_threshold")
|
||||||
|
and num_classes == 2 and y_val.size >= 2):
|
||||||
|
val_threshold = tune_binary_threshold(y_val, p_val[:, 1])
|
||||||
|
|
||||||
|
y_te = p_te = ids_te = z_te = None
|
||||||
|
test_auc = test_acc = test_n = nan
|
||||||
|
if ctx["test_loader"] is not None:
|
||||||
|
y_te, p_te, ids_te, z_te = collect_probs(bridge, primary_head, ctx["sc"], towers,
|
||||||
|
stage_models, cfg_stages, ctx["test_loader"],
|
||||||
|
device, num_classes)
|
||||||
|
test_acc, test_auc, test_n = (score_arrays(y_te, p_te, num_classes)
|
||||||
|
if y_te.size else (nan, nan, nan))
|
||||||
|
|
||||||
|
all_metrics.update({
|
||||||
|
f"{name}_val_auc": val_auc,
|
||||||
|
f"{name}_val_acc": val_acc,
|
||||||
|
f"{name}_val_n": val_n,
|
||||||
|
f"{name}_val_kappa": ext.get("kappa", nan),
|
||||||
|
f"{name}_val_mcc": ext.get("mcc", nan),
|
||||||
|
f"{name}_val_f1": ext.get("macro_f1", nan),
|
||||||
|
f"{name}_val_threshold": val_threshold,
|
||||||
|
f"{name}_test_auc": test_auc,
|
||||||
|
f"{name}_test_acc": test_acc,
|
||||||
|
f"{name}_test_n": test_n,
|
||||||
|
})
|
||||||
|
all_preds[name] = {
|
||||||
|
"val_y": y_val, "val_p": p_val, "val_ids": ids_val, "val_z": z_val,
|
||||||
|
"test_y": y_te, "test_p": p_te, "test_ids": ids_te, "test_z": z_te,
|
||||||
|
}
|
||||||
|
|
||||||
|
return updated, all_metrics, all_preds
|
||||||
@@ -1,11 +1,13 @@
|
|||||||
"""stages/warm — warm stage runner: pre-trains a single tower with a temporary probe."""
|
"""stages/warm — warm stage runner: pre-trains a single tower and optional real head."""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import importlib
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
from v4.classes.dataset import to_label_tensor
|
from v4.classes.dataset import to_label_tensor
|
||||||
from v4.classes.stages.helpers import phase_for_epoch
|
from v4.classes.stages.helpers import class_weights_from_shell
|
||||||
|
|
||||||
|
|
||||||
def run(
|
def run(
|
||||||
@@ -20,16 +22,27 @@ def run(
|
|||||||
fold: int,
|
fold: int,
|
||||||
_make_loader,
|
_make_loader,
|
||||||
_balanced_sampler,
|
_balanced_sampler,
|
||||||
) -> None:
|
stage_models: dict | None = None,
|
||||||
"""Pre-train one tower using a temporary linear probe (probe discarded after)."""
|
cfg_stages: list[dict] | None = None,
|
||||||
|
) -> dict:
|
||||||
|
"""Pre-train one tower.
|
||||||
|
|
||||||
|
If ``head_name`` is set on the stage config, train that real downstream head
|
||||||
|
and return it in ``stage_models``. Otherwise, fall back to a temporary linear
|
||||||
|
probe for backward-compatible representation warmup.
|
||||||
|
"""
|
||||||
tower_name = stage_cfg["tower"]
|
tower_name = stage_cfg["tower"]
|
||||||
n_epochs = stage_cfg.get("epochs", 0)
|
n_epochs = stage_cfg.get("epochs", 0)
|
||||||
level = stage_cfg["level"]
|
level = stage_cfg["level"]
|
||||||
|
shell_filter = stage_cfg.get("shell_filter", {})
|
||||||
|
stage_models = dict(stage_models or {})
|
||||||
|
cfg_stages = list(cfg_stages or [])
|
||||||
|
|
||||||
if n_epochs == 0:
|
if n_epochs == 0:
|
||||||
return
|
return stage_models
|
||||||
|
|
||||||
s_train = data.build_shells(split.train, level=level, label_filter=label_filter)
|
s_train = data.build_shells(split.train, level=level, label_filter=label_filter,
|
||||||
|
**shell_filter)
|
||||||
bs = cfg["training"]["batch_size"]
|
bs = cfg["training"]["batch_size"]
|
||||||
loader = _make_loader(
|
loader = _make_loader(
|
||||||
s_train, {tower_name: towers[tower_name]},
|
s_train, {tower_name: towers[tower_name]},
|
||||||
@@ -41,13 +54,32 @@ def run(
|
|||||||
for p in t.parameters():
|
for p in t.parameters():
|
||||||
p.requires_grad_(n == tower_name)
|
p.requires_grad_(n == tower_name)
|
||||||
|
|
||||||
|
head_name = stage_cfg.get("head_name")
|
||||||
|
if head_name:
|
||||||
|
head_cfg = next((s for s in cfg_stages if s.get("name") == head_name), None)
|
||||||
|
if head_cfg is None:
|
||||||
|
raise ValueError(f"warm stage requested head_name={head_name!r}, but no such head exists")
|
||||||
|
h_mod = importlib.import_module(head_cfg.get("module", "v4.classes.heads.classifier"))
|
||||||
|
h_cls = getattr(h_mod, head_cfg.get("class", "ClassificationHead"))
|
||||||
|
probe = stage_models.get(head_name)
|
||||||
|
if probe is None:
|
||||||
|
probe = h_cls(towers[tower_name].out_dim, num_classes, **head_cfg.get("args", {}))
|
||||||
|
probe = probe.to(device)
|
||||||
|
else:
|
||||||
probe = torch.nn.Linear(towers[tower_name].out_dim, num_classes).to(device)
|
probe = torch.nn.Linear(towers[tower_name].out_dim, num_classes).to(device)
|
||||||
|
|
||||||
opt = torch.optim.Adam(
|
opt = torch.optim.Adam(
|
||||||
list(towers[tower_name].parameters()) + list(probe.parameters()),
|
list(towers[tower_name].parameters()) + list(probe.parameters()),
|
||||||
lr=cfg["training"]["lr"],
|
lr=cfg["training"]["lr"],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
cw = class_weights_from_shell(
|
||||||
|
s_train, num_classes, device,
|
||||||
|
enabled=cfg["training"].get("class_weighted", False),
|
||||||
|
)
|
||||||
|
|
||||||
towers[tower_name].train()
|
towers[tower_name].train()
|
||||||
|
probe.train()
|
||||||
for epoch in range(n_epochs):
|
for epoch in range(n_epochs):
|
||||||
total_loss = total_correct = total_n = 0
|
total_loss = total_correct = total_n = 0
|
||||||
for batch in loader:
|
for batch in loader:
|
||||||
@@ -57,7 +89,7 @@ def run(
|
|||||||
continue
|
continue
|
||||||
y_t = to_label_tensor(y, device)
|
y_t = to_label_tensor(y, device)
|
||||||
logits = probe(towers[tower_name](x.to(device)))
|
logits = probe(towers[tower_name](x.to(device)))
|
||||||
loss = F.cross_entropy(logits, y_t)
|
loss = F.cross_entropy(logits, y_t, weight=cw)
|
||||||
opt.zero_grad(); loss.backward(); opt.step()
|
opt.zero_grad(); loss.backward(); opt.step()
|
||||||
total_loss += loss.item() * len(y_t)
|
total_loss += loss.item() * len(y_t)
|
||||||
total_correct += int((logits.argmax(1) == y_t).sum())
|
total_correct += int((logits.argmax(1) == y_t).sum())
|
||||||
@@ -71,3 +103,7 @@ def run(
|
|||||||
for t in towers.values():
|
for t in towers.values():
|
||||||
for p in t.parameters():
|
for p in t.parameters():
|
||||||
p.requires_grad_(True)
|
p.requires_grad_(True)
|
||||||
|
|
||||||
|
if head_name:
|
||||||
|
stage_models[head_name] = probe
|
||||||
|
return stage_models
|
||||||
|
|||||||
@@ -2,6 +2,25 @@
|
|||||||
|
|
||||||
Self-contained: no v3 dependencies.
|
Self-contained: no v3 dependencies.
|
||||||
Inherits get_sample dispatch from TowerBase.
|
Inherits get_sample dispatch from TowerBase.
|
||||||
|
|
||||||
|
Geometry injection (EPC consumption)
|
||||||
|
--------------------------------------
|
||||||
|
When geom_dim > 0, ClinicalEncoder requests the "geometry_vectors" key from EPC
|
||||||
|
during early_pass and appends the geometry features to every clinical vector.
|
||||||
|
The input layer is sized to clinical_data.feature_dim + geom_dim automatically.
|
||||||
|
|
||||||
|
Config example (cd tower consuming geometry):
|
||||||
|
{
|
||||||
|
"name": "cd",
|
||||||
|
"module": "v4.classes.towers.clinical_tower",
|
||||||
|
"class": "ClinicalEncoder",
|
||||||
|
"data_source": "matrix",
|
||||||
|
"epc_requests": ["geometry_vectors"],
|
||||||
|
"args": {
|
||||||
|
"hidden_dim": 128,
|
||||||
|
"geom_dim": 5
|
||||||
|
}
|
||||||
|
}
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -14,7 +33,7 @@ from v4.classes.accessory.se_block import SEBlock
|
|||||||
|
|
||||||
|
|
||||||
class ClinicalEncoder(TowerBase):
|
class ClinicalEncoder(TowerBase):
|
||||||
"""MLP over tabular clinical features.
|
"""MLP over tabular clinical features, with optional geometry vector injection.
|
||||||
|
|
||||||
clinical_data : ClinicalDataView — provides feature_dim, vectorize_entity, side_map
|
clinical_data : ClinicalDataView — provides feature_dim, vectorize_entity, side_map
|
||||||
hidden_dim : output embedding dimensionality
|
hidden_dim : output embedding dimensionality
|
||||||
@@ -22,8 +41,12 @@ class ClinicalEncoder(TowerBase):
|
|||||||
use_se : wrap output with SEBlock channel gating
|
use_se : wrap output with SEBlock channel gating
|
||||||
se_reduction : SEBlock bottleneck factor
|
se_reduction : SEBlock bottleneck factor
|
||||||
se_pre_norm : apply LayerNorm before SEBlock
|
se_pre_norm : apply LayerNorm before SEBlock
|
||||||
|
geom_dim : number of geometry features to append from EPC (0 = disabled)
|
||||||
|
requires epc_requests: ["geometry_vectors"] in tower config
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
EPC_GEOMETRY_KEY = "geometry_vectors"
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
clinical_data,
|
clinical_data,
|
||||||
@@ -32,11 +55,14 @@ class ClinicalEncoder(TowerBase):
|
|||||||
use_se: bool = False,
|
use_se: bool = False,
|
||||||
se_reduction: int = 16,
|
se_reduction: int = 16,
|
||||||
se_pre_norm: bool = True,
|
se_pre_norm: bool = True,
|
||||||
|
geom_dim: int = 0,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.clinical_data = clinical_data
|
self.clinical_data = clinical_data
|
||||||
self._out_dim = hidden_dim
|
self._out_dim = hidden_dim
|
||||||
feature_dim = clinical_data.feature_dim
|
self._geom_dim = geom_dim
|
||||||
|
self._geom_vectors: dict | None = None # filled by early_pass when geom_dim > 0
|
||||||
|
feature_dim = clinical_data.feature_dim + geom_dim
|
||||||
|
|
||||||
self.block0 = nn.Sequential(
|
self.block0 = nn.Sequential(
|
||||||
nn.Linear(feature_dim, hidden_dim),
|
nn.Linear(feature_dim, hidden_dim),
|
||||||
@@ -53,6 +79,12 @@ class ClinicalEncoder(TowerBase):
|
|||||||
self.tower_ln = nn.LayerNorm(hidden_dim) if se_pre_norm else nn.Identity()
|
self.tower_ln = nn.LayerNorm(hidden_dim) if se_pre_norm else nn.Identity()
|
||||||
self.tower_se = SEBlock(hidden_dim, reduction=se_reduction, residual=True) if use_se else None
|
self.tower_se = SEBlock(hidden_dim, reduction=se_reduction, residual=True) if use_se else None
|
||||||
|
|
||||||
|
# ── EPC early_pass ───────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def early_pass(self, context) -> None:
|
||||||
|
if self._geom_dim > 0:
|
||||||
|
self._geom_vectors = context.require(self.EPC_GEOMETRY_KEY)
|
||||||
|
|
||||||
# ── TowerBase interface ──────────────────────────────────────────────────
|
# ── TowerBase interface ──────────────────────────────────────────────────
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -65,6 +97,14 @@ class ClinicalEncoder(TowerBase):
|
|||||||
|
|
||||||
def _get(self, *ids) -> torch.Tensor:
|
def _get(self, *ids) -> torch.Tensor:
|
||||||
arr = self.clinical_data.vectorize_entity(*ids)
|
arr = self.clinical_data.vectorize_entity(*ids)
|
||||||
|
if self._geom_dim > 0 and self._geom_vectors is not None:
|
||||||
|
pid = int(ids[0])
|
||||||
|
eye = str(ids[1]) if len(ids) > 1 else "OD"
|
||||||
|
geom = self._geom_vectors.get(
|
||||||
|
(pid, eye),
|
||||||
|
np.zeros(self._geom_dim, dtype=np.float32),
|
||||||
|
)
|
||||||
|
arr = np.concatenate([arr, geom[: self._geom_dim]])
|
||||||
return torch.from_numpy(arr.astype(np.float32, copy=False))
|
return torch.from_numpy(arr.astype(np.float32, copy=False))
|
||||||
|
|
||||||
# ── nn.Module forward ────────────────────────────────────────────────────
|
# ── nn.Module forward ────────────────────────────────────────────────────
|
||||||
|
|||||||
@@ -0,0 +1,237 @@
|
|||||||
|
"""geometry_tower — GeometrySegEncoder for v4.
|
||||||
|
|
||||||
|
A CNN tower that takes a per-eye disc/cup *segmentation map* as input (rather
|
||||||
|
than the raw fundus image) and contributes its pooled embedding to fusion.
|
||||||
|
|
||||||
|
Seg maps are produced by an underlying loader (GT contour rasterisation or
|
||||||
|
UNet inference) during early_pass, then cached per fold.
|
||||||
|
|
||||||
|
UNet fine-tuning lives in early_pass too — the loader's `finetune(train_samples)`
|
||||||
|
call uses only the training split, then precompute() runs inference on all
|
||||||
|
fold samples (train + val + test).
|
||||||
|
|
||||||
|
Config example:
|
||||||
|
{
|
||||||
|
"name": "geom",
|
||||||
|
"module": "v4.classes.towers.geometry_tower",
|
||||||
|
"class": "GeometrySegEncoder",
|
||||||
|
"data_source": "image",
|
||||||
|
"args": {
|
||||||
|
"backbone": "resnet18",
|
||||||
|
"channels": 3,
|
||||||
|
"target_size": 224,
|
||||||
|
"augment": true,
|
||||||
|
"seg_source": "gt",
|
||||||
|
"contour_dir": "Papila/ExpertsSegmentations/Contours"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import math
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
from torchvision import models
|
||||||
|
|
||||||
|
_REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||||
|
if str(_REPO_ROOT) not in sys.path:
|
||||||
|
sys.path.insert(0, str(_REPO_ROOT))
|
||||||
|
|
||||||
|
from v4.classes.accessory.backbones import build_backbone
|
||||||
|
from v4.classes.towerbase import TowerBase
|
||||||
|
|
||||||
|
|
||||||
|
class GeometrySegEncoder(TowerBase):
|
||||||
|
"""CNN tower over disc/cup segmentation maps.
|
||||||
|
|
||||||
|
image_data : ImageDataView — provides get_image_path(*ids) and side_map.
|
||||||
|
Must implement build_seg_map_loader(source, **kwargs).
|
||||||
|
backbone : backbone key (see accessory/backbones.py)
|
||||||
|
channels : 1 (label map in [0,1]) or 3 (one-hot bg/rim/cup)
|
||||||
|
target_size : CNN input spatial size (cached arrays already at this size)
|
||||||
|
augment : random flip + 90° rotation at training time
|
||||||
|
freeze_ratio : fraction of early backbone blocks to freeze in [0, 1]
|
||||||
|
seg_source : passed to image_data.build_seg_map_loader (e.g. "gt", "unet")
|
||||||
|
**seg_kwargs : forwarded to build_seg_map_loader
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
image_data,
|
||||||
|
backbone: str = "resnet18",
|
||||||
|
channels: int = 3,
|
||||||
|
target_size: int = 224,
|
||||||
|
augment: bool = True,
|
||||||
|
freeze_ratio: float = 0.0,
|
||||||
|
seg_source: str = "gt",
|
||||||
|
**seg_kwargs: Any,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.image_data = image_data
|
||||||
|
self._channels = channels
|
||||||
|
self._target_size = target_size
|
||||||
|
self._augment = augment
|
||||||
|
|
||||||
|
if not hasattr(image_data, "build_seg_map_loader"):
|
||||||
|
raise TypeError(
|
||||||
|
f"GeometrySegEncoder requires image_data to implement "
|
||||||
|
f"build_seg_map_loader(), but {type(image_data).__name__} does not."
|
||||||
|
)
|
||||||
|
loader_kwargs = {
|
||||||
|
"channels": channels,
|
||||||
|
"target_size": target_size,
|
||||||
|
**seg_kwargs,
|
||||||
|
}
|
||||||
|
self._loader = image_data.build_seg_map_loader(seg_source, **loader_kwargs)
|
||||||
|
self._seg_cache: dict = {}
|
||||||
|
self._seg_source = seg_source
|
||||||
|
|
||||||
|
self.backbone, self._base_dim, self._blocks = build_backbone(backbone, freeze_ratio)
|
||||||
|
if channels != 3:
|
||||||
|
self._adapt_first_conv(channels)
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"[GeometrySegEncoder] backbone={backbone} channels={channels} "
|
||||||
|
f"target_size={target_size} seg_source={seg_source}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── TowerBase interface ──────────────────────────────────────────────────
|
||||||
|
|
||||||
|
@property
|
||||||
|
def out_dim(self) -> int:
|
||||||
|
return self._base_dim
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _side_map(self) -> dict[str, str]:
|
||||||
|
return self.image_data.side_map
|
||||||
|
|
||||||
|
def _get(self, *ids) -> torch.Tensor:
|
||||||
|
key = tuple(ids)
|
||||||
|
arr = self._seg_cache.get(key)
|
||||||
|
if arr is None:
|
||||||
|
arr = np.zeros(
|
||||||
|
(self._channels, self._target_size, self._target_size),
|
||||||
|
dtype=np.float32,
|
||||||
|
)
|
||||||
|
if self.training and self._augment:
|
||||||
|
arr = self._augment_array(arr)
|
||||||
|
return torch.from_numpy(np.ascontiguousarray(arr))
|
||||||
|
|
||||||
|
# ── EPC early_pass ───────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def early_pass(self, context) -> None:
|
||||||
|
data = context.require("data")
|
||||||
|
split = context.require("split")
|
||||||
|
|
||||||
|
train_samples = self._collect_samples(split.train, data)
|
||||||
|
all_samples = self._collect_samples(split.train, data)
|
||||||
|
all_samples += self._collect_samples(split.val, data)
|
||||||
|
if split.test is not None:
|
||||||
|
all_samples += self._collect_samples(split.test, data)
|
||||||
|
|
||||||
|
# Reset per-fold state if loader supports it (UNet only).
|
||||||
|
if hasattr(self._loader, "reset_cache"):
|
||||||
|
self._loader.reset_cache()
|
||||||
|
if hasattr(self._loader, "reset_weights"):
|
||||||
|
self._loader.reset_weights()
|
||||||
|
if hasattr(self._loader, "finetune"):
|
||||||
|
self._loader.finetune(train_samples)
|
||||||
|
|
||||||
|
self._loader.precompute(all_samples)
|
||||||
|
self._seg_cache = self._loader.all_seg_maps()
|
||||||
|
print(
|
||||||
|
f"[GeometrySegEncoder] cached {len(self._seg_cache)} seg maps for fold",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── nn.Module forward ────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
y = self.backbone(x)
|
||||||
|
if y.dim() > 2:
|
||||||
|
y = y.flatten(1)
|
||||||
|
return y
|
||||||
|
|
||||||
|
# ── Utilities ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def set_freeze_ratio(self, ratio: float) -> None:
|
||||||
|
r = max(0.0, min(1.0, float(ratio)))
|
||||||
|
n_freeze = int(math.floor(len(self._blocks) * r))
|
||||||
|
for b in self._blocks:
|
||||||
|
for p in b.parameters():
|
||||||
|
p.requires_grad = True
|
||||||
|
for b in self._blocks[:n_freeze]:
|
||||||
|
for p in b.parameters():
|
||||||
|
p.requires_grad = False
|
||||||
|
|
||||||
|
# ── Internals ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _collect_samples(self, df, data) -> list:
|
||||||
|
"""Build (pid, eye, image_path) tuples from a split DataFrame."""
|
||||||
|
if df is None or len(df) == 0:
|
||||||
|
return []
|
||||||
|
pc = data.patient_col
|
||||||
|
out = []
|
||||||
|
for _, row in df.iterrows():
|
||||||
|
pid = int(row[pc])
|
||||||
|
eye = str(row.get("eyeID", "OD"))
|
||||||
|
out.append((pid, eye, data.image.get_image_path(pid, eye)))
|
||||||
|
return out
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _augment_array(arr: np.ndarray) -> np.ndarray:
|
||||||
|
"""Random flip + 90° rotation on a (C, H, W) seg-map array."""
|
||||||
|
if np.random.rand() < 0.5:
|
||||||
|
arr = arr[:, :, ::-1]
|
||||||
|
if np.random.rand() < 0.5:
|
||||||
|
arr = arr[:, ::-1, :]
|
||||||
|
k = int(np.random.randint(0, 4))
|
||||||
|
if k:
|
||||||
|
arr = np.rot90(arr, k=k, axes=(1, 2))
|
||||||
|
return arr
|
||||||
|
|
||||||
|
def _adapt_first_conv(self, in_channels: int) -> None:
|
||||||
|
"""Replace the first Conv2d to accept a non-3-channel input.
|
||||||
|
|
||||||
|
Pretrained weights are averaged across the original input channels and
|
||||||
|
broadcast across the new ones.
|
||||||
|
"""
|
||||||
|
first = self._find_first_conv(self.backbone)
|
||||||
|
new = nn.Conv2d(
|
||||||
|
in_channels,
|
||||||
|
first.out_channels,
|
||||||
|
kernel_size=first.kernel_size,
|
||||||
|
stride=first.stride,
|
||||||
|
padding=first.padding,
|
||||||
|
bias=first.bias is not None,
|
||||||
|
)
|
||||||
|
with torch.no_grad():
|
||||||
|
new.weight.copy_(
|
||||||
|
first.weight.mean(dim=1, keepdim=True).expand_as(new.weight)
|
||||||
|
)
|
||||||
|
if first.bias is not None:
|
||||||
|
new.bias.copy_(first.bias)
|
||||||
|
self._replace_first_conv(self.backbone, new)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _find_first_conv(module: nn.Module) -> nn.Conv2d:
|
||||||
|
for m in module.modules():
|
||||||
|
if isinstance(m, nn.Conv2d):
|
||||||
|
return m
|
||||||
|
raise RuntimeError("No Conv2d found in backbone")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _replace_first_conv(cls, module: nn.Module, new_conv: nn.Conv2d) -> bool:
|
||||||
|
for name, child in module.named_children():
|
||||||
|
if isinstance(child, nn.Conv2d):
|
||||||
|
setattr(module, name, new_conv)
|
||||||
|
return True
|
||||||
|
if cls._replace_first_conv(child, new_conv):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
@@ -2,30 +2,83 @@
|
|||||||
|
|
||||||
Self-contained: no v3 dependencies.
|
Self-contained: no v3 dependencies.
|
||||||
Inherits get_sample dispatch from TowerBase.
|
Inherits get_sample dispatch from TowerBase.
|
||||||
|
|
||||||
|
Geometry injection (EPC supply)
|
||||||
|
--------------------------------
|
||||||
|
When geometry_source is set, ImageEncoder asks the image_data view for a loader
|
||||||
|
via image_data.build_geometry_loader(source, **kwargs). The view is responsible
|
||||||
|
for understanding what that source means for its specific domain (fundus contours,
|
||||||
|
U-Net segmentations, cat ear landmarks, etc.).
|
||||||
|
|
||||||
|
During early_pass the loader pre-computes all per-entity geometry vectors and
|
||||||
|
publishes them to the EarlyPassContext under the key "geometry_vectors"
|
||||||
|
({(entity_id...): np.ndarray of length geom_dim}). ClinicalEncoder (or any
|
||||||
|
other tower with epc_requests: ["geometry_vectors"]) can then consume them.
|
||||||
|
|
||||||
|
The tower reads feature_dim and feature_names from the loader instance, so it
|
||||||
|
can log geometry info without knowing anything about CDR, disc masks, or other
|
||||||
|
domain-specific concepts.
|
||||||
|
|
||||||
|
Config example:
|
||||||
|
{
|
||||||
|
"name": "img",
|
||||||
|
"module": "v4.classes.towers.image_tower",
|
||||||
|
"class": "ImageEncoder",
|
||||||
|
"data_source": "image",
|
||||||
|
"epc_supplies": ["geometry_vectors"],
|
||||||
|
"args": {
|
||||||
|
"backbone": "refugelike",
|
||||||
|
"augment": true,
|
||||||
|
"geometry_source": "gt",
|
||||||
|
"contour_dir": "Papila/ExpertsSegmentations/Contours"
|
||||||
|
}
|
||||||
|
}
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import math
|
import math
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
|
_REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||||
|
if str(_REPO_ROOT) not in sys.path:
|
||||||
|
sys.path.insert(0, str(_REPO_ROOT))
|
||||||
|
|
||||||
from v4.classes.towerbase import TowerBase
|
from v4.classes.towerbase import TowerBase
|
||||||
from v4.classes.accessory.backbones import build_backbone
|
from v4.classes.accessory.backbones import build_backbone
|
||||||
from v4.classes.accessory.se_block import SEBlock
|
from v4.classes.accessory.se_block import SEBlock
|
||||||
from v4.classes.accessory.transforms import build_backbone_transform, build_eval_transform
|
from v4.classes.accessory.transforms import (
|
||||||
|
build_backbone_transform, build_eval_transform, build_split_transforms,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ImageEncoder(TowerBase):
|
class ImageEncoder(TowerBase):
|
||||||
"""Vision backbone → pooled feature vector.
|
"""Vision backbone → pooled feature vector.
|
||||||
|
|
||||||
image_data : ImageDataView — provides load_image(*ids) and side_map
|
image_data : ImageDataView — provides load_image(*ids) and side_map.
|
||||||
|
Must implement build_geometry_loader(source, **kwargs)
|
||||||
|
if geometry_source is set.
|
||||||
backbone : backbone key (see accessory/backbones.py)
|
backbone : backbone key (see accessory/backbones.py)
|
||||||
freeze_ratio : fraction of early blocks to freeze in [0, 1]
|
freeze_ratio : fraction of early blocks to freeze in [0, 1]
|
||||||
use_se : apply SE attention over the pooled feature vector
|
use_se : apply SE attention over the pooled feature vector
|
||||||
augment : include random flip/rotation/jitter in the train transform
|
augment : include random flip/rotation/jitter in the train transform
|
||||||
|
cache_transformed : if True, cache resized + ToTensor'd float32 [0, 1] CHW
|
||||||
|
tensors per fold. Per-batch cost drops to augment +
|
||||||
|
Normalize on tensors only (no PIL, no Resize, no decode).
|
||||||
|
Memory: ~3 × crop_size² × 4B per cached image.
|
||||||
|
Cache is rebuilt at the start of every fold via early_pass.
|
||||||
|
geometry_source : source key passed to image_data.build_geometry_loader()
|
||||||
|
(e.g. "gt", "unet"). None = geometry disabled.
|
||||||
|
**geom_kwargs : forwarded verbatim to build_geometry_loader() — e.g.
|
||||||
|
contour_dir="Papila/ExpertsSegmentations/Contours"
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
EPC_GEOMETRY_KEY = "geometry_vectors"
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
image_data,
|
image_data,
|
||||||
@@ -35,17 +88,42 @@ class ImageEncoder(TowerBase):
|
|||||||
se_reduction: int = 16,
|
se_reduction: int = 16,
|
||||||
se_pre_norm: bool = True,
|
se_pre_norm: bool = True,
|
||||||
augment: bool = True,
|
augment: bool = True,
|
||||||
|
cache_transformed: bool = False,
|
||||||
|
geometry_source: str | None = None,
|
||||||
|
**geom_kwargs: Any,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.image_data = image_data
|
self.image_data = image_data
|
||||||
self._name = backbone
|
self._name = backbone
|
||||||
self.backbone, self._base_dim, self._blocks = build_backbone(backbone, freeze_ratio)
|
self.backbone, self._base_dim, self._blocks = build_backbone(backbone, freeze_ratio)
|
||||||
|
|
||||||
|
self._cache_transformed = cache_transformed
|
||||||
|
if cache_transformed:
|
||||||
|
self._precache_tf, self._post_train_tf = build_split_transforms(backbone, augment=augment)
|
||||||
|
_, self._post_eval_tf = build_split_transforms(backbone, augment=False)
|
||||||
|
self._tensor_cache: dict[tuple, torch.Tensor] = {}
|
||||||
|
else:
|
||||||
self.transform = build_backbone_transform(backbone, augment=augment)
|
self.transform = build_backbone_transform(backbone, augment=augment)
|
||||||
self.eval_transform = build_eval_transform(backbone)
|
self.eval_transform = build_eval_transform(backbone)
|
||||||
|
|
||||||
self.tower_ln = nn.LayerNorm(self._base_dim) if se_pre_norm else nn.Identity()
|
self.tower_ln = nn.LayerNorm(self._base_dim) if se_pre_norm else nn.Identity()
|
||||||
self.tower_se = SEBlock(self._base_dim, reduction=se_reduction, residual=True) if use_se else None
|
self.tower_se = SEBlock(self._base_dim, reduction=se_reduction, residual=True) if use_se else None
|
||||||
|
|
||||||
|
self._geom_loader = None
|
||||||
|
if geometry_source is not None:
|
||||||
|
if not hasattr(image_data, "build_geometry_loader"):
|
||||||
|
raise TypeError(
|
||||||
|
f"ImageEncoder geometry_source={geometry_source!r} requires "
|
||||||
|
f"image_data to implement build_geometry_loader(), "
|
||||||
|
f"but {type(image_data).__name__} does not."
|
||||||
|
)
|
||||||
|
self._geom_loader = image_data.build_geometry_loader(geometry_source, **geom_kwargs)
|
||||||
|
print(
|
||||||
|
f"[ImageEncoder] geometry_source={geometry_source!r} "
|
||||||
|
f"features={self._geom_loader.feature_names}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
# ── TowerBase interface ──────────────────────────────────────────────────
|
# ── TowerBase interface ──────────────────────────────────────────────────
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -57,10 +135,61 @@ class ImageEncoder(TowerBase):
|
|||||||
return self.image_data.side_map
|
return self.image_data.side_map
|
||||||
|
|
||||||
def _get(self, *ids) -> torch.Tensor:
|
def _get(self, *ids) -> torch.Tensor:
|
||||||
|
if self._cache_transformed:
|
||||||
|
key = tuple(ids)
|
||||||
|
cached = self._tensor_cache.get(key)
|
||||||
|
if cached is None:
|
||||||
|
cached = self._precache_tf(self.image_data.load_image(*ids))
|
||||||
|
self._tensor_cache[key] = cached
|
||||||
|
tail = self._post_train_tf if self.training else self._post_eval_tf
|
||||||
|
return tail(cached)
|
||||||
img = self.image_data.load_image(*ids)
|
img = self.image_data.load_image(*ids)
|
||||||
t = self.transform if self.training else self.eval_transform
|
t = self.transform if self.training else self.eval_transform
|
||||||
return t(img)
|
return t(img)
|
||||||
|
|
||||||
|
# ── EPC early_pass ───────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def early_pass(self, context) -> None:
|
||||||
|
"""Per-fold setup: warm tensor cache (if enabled), publish geometry vectors."""
|
||||||
|
data = context.require("data")
|
||||||
|
|
||||||
|
if self._cache_transformed:
|
||||||
|
self._tensor_cache.clear()
|
||||||
|
n = self._warm_tensor_cache(data, context.require("split"))
|
||||||
|
print(
|
||||||
|
f"[ImageEncoder] warmed transformed-tensor cache for {n} entries "
|
||||||
|
f"({self._name})",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self._geom_loader is None:
|
||||||
|
return
|
||||||
|
self._geom_loader.precompute(data.df, patient_col=data.patient_col)
|
||||||
|
vecs = self._geom_loader.all_vectors()
|
||||||
|
context.put(self.EPC_GEOMETRY_KEY, vecs)
|
||||||
|
print(
|
||||||
|
f"[ImageEncoder] published {len(vecs)} geometry vectors "
|
||||||
|
f"(dim={self._geom_loader.feature_dim}) to EPC key '{self.EPC_GEOMETRY_KEY}'",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _warm_tensor_cache(self, data, split) -> int:
|
||||||
|
"""Pre-fill the per-tower tensor cache for all entries in this fold's splits."""
|
||||||
|
seen: set[tuple] = set()
|
||||||
|
for df in (split.train, split.val, split.test):
|
||||||
|
if df is None or len(df) == 0:
|
||||||
|
continue
|
||||||
|
pc = data.patient_col
|
||||||
|
for _, row in df.iterrows():
|
||||||
|
pid = int(row[pc])
|
||||||
|
eye = str(row.get("eyeID", "OD"))
|
||||||
|
key = (pid, eye)
|
||||||
|
if key in self._tensor_cache or key in seen:
|
||||||
|
continue
|
||||||
|
self._tensor_cache[key] = self._precache_tf(self.image_data.load_image(pid, eye))
|
||||||
|
seen.add(key)
|
||||||
|
return len(self._tensor_cache)
|
||||||
|
|
||||||
# ── nn.Module forward ────────────────────────────────────────────────────
|
# ── nn.Module forward ────────────────────────────────────────────────────
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
|||||||
+110
-7
@@ -29,7 +29,8 @@ sys.path.insert(0, str(REPO_ROOT))
|
|||||||
from v4.classes.dataset import LoaderShell, HTDataset, ht_collate
|
from v4.classes.dataset import LoaderShell, HTDataset, ht_collate
|
||||||
from v4.classes.utils import seed_everything, choose_device
|
from v4.classes.utils import seed_everything, choose_device
|
||||||
from v4.classes.split_manager import SplitManager
|
from v4.classes.split_manager import SplitManager
|
||||||
from v4.classes.stages import warm, fusion
|
from v4.classes.stages import warm, fusion, parallel
|
||||||
|
from v4.classes.logging.prediction_store import PredictionStore, FeatureStore
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -167,28 +168,49 @@ def run_fold(fold: int, splits, cfg: dict, data, num_classes: int, device) -> di
|
|||||||
if hasattr(tower, "early_pass"):
|
if hasattr(tower, "early_pass"):
|
||||||
tower.early_pass(context)
|
tower.early_pass(context)
|
||||||
|
|
||||||
|
# Flatten parallel wrappers so sub-stage configs are addressable by name.
|
||||||
|
flat_stages: list[dict] = []
|
||||||
|
for s in cfg_stages:
|
||||||
|
if s.get("type") == "parallel":
|
||||||
|
flat_stages.extend(s["stages"])
|
||||||
|
else:
|
||||||
|
flat_stages.append(s)
|
||||||
|
|
||||||
stage_models: dict = {}
|
stage_models: dict = {}
|
||||||
fold_result = {"fold": fold}
|
fold_result = {"fold": fold}
|
||||||
|
fold_preds: dict = {} # stage_name → pred_data
|
||||||
|
|
||||||
for stage_cfg in cfg_stages:
|
for stage_cfg in cfg_stages:
|
||||||
stype = stage_cfg["type"]
|
stype = stage_cfg["type"]
|
||||||
|
|
||||||
if stype == "warm":
|
if stype == "warm":
|
||||||
warm.run(stage_cfg, towers, data, split, label_filter,
|
stage_models = warm.run(
|
||||||
|
stage_cfg, towers, data, split, label_filter,
|
||||||
cfg, num_classes, device, fold,
|
cfg, num_classes, device, fold,
|
||||||
_make_loader, _balanced_sampler)
|
_make_loader, _balanced_sampler, stage_models, flat_stages,
|
||||||
|
)
|
||||||
|
|
||||||
elif stype == "fusion":
|
elif stype == "fusion":
|
||||||
stage_models, metrics = fusion.run(
|
stage_models, metrics, preds = fusion.run(
|
||||||
stage_cfg, cfg, towers, stage_models, data, split,
|
stage_cfg, cfg, towers, stage_models, data, split,
|
||||||
label_filter, num_classes, device, fold, cfg_stages,
|
label_filter, num_classes, device, fold, flat_stages,
|
||||||
_make_loader,
|
_make_loader,
|
||||||
)
|
)
|
||||||
fold_result.update(metrics)
|
fold_result.update(metrics)
|
||||||
|
fold_preds.update(preds)
|
||||||
|
|
||||||
|
elif stype == "parallel":
|
||||||
|
stage_models, metrics, preds = parallel.run(
|
||||||
|
stage_cfg, cfg, towers, stage_models, data, split,
|
||||||
|
label_filter, num_classes, device, fold, flat_stages,
|
||||||
|
_make_loader, _balanced_sampler,
|
||||||
|
)
|
||||||
|
fold_result.update(metrics)
|
||||||
|
fold_preds.update(preds)
|
||||||
|
|
||||||
# head stages are handled inside fusion.run
|
# head stages are handled inside fusion.run
|
||||||
|
|
||||||
return fold_result
|
return fold_result, fold_preds
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -240,7 +262,10 @@ def main():
|
|||||||
out_dir.mkdir(parents=True, exist_ok=True)
|
out_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
eval_stage = cfg.get("eval_stage", "hb")
|
eval_stage = cfg.get("eval_stage", "hb")
|
||||||
|
save_predictions = cfg.get("save_predictions", False)
|
||||||
|
save_features = cfg.get("save_features", False)
|
||||||
fold_results = []
|
fold_results = []
|
||||||
|
eval_stage_preds = [] # list[dict] — one per fold, only for eval_stage
|
||||||
t0 = time.time()
|
t0 = time.time()
|
||||||
|
|
||||||
for fold in range(cfg.get("folds", 5)):
|
for fold in range(cfg.get("folds", 5)):
|
||||||
@@ -248,8 +273,10 @@ def main():
|
|||||||
n_train = split.train[group_col].nunique() if group_col else len(split.train)
|
n_train = split.train[group_col].nunique() if group_col else len(split.train)
|
||||||
print(f"\n── fold {fold+1}/{cfg.get('folds', 5)} train_groups={n_train} ──",
|
print(f"\n── fold {fold+1}/{cfg.get('folds', 5)} train_groups={n_train} ──",
|
||||||
flush=True)
|
flush=True)
|
||||||
result = run_fold(fold, splits, cfg, data, num_classes, device)
|
result, fold_preds = run_fold(fold, splits, cfg, data, num_classes, device)
|
||||||
fold_results.append(result)
|
fold_results.append(result)
|
||||||
|
if save_predictions and eval_stage in fold_preds:
|
||||||
|
eval_stage_preds.append(fold_preds[eval_stage])
|
||||||
print(
|
print(
|
||||||
f" fold{fold+1} DONE"
|
f" fold{fold+1} DONE"
|
||||||
f" val_auc={result.get(f'{eval_stage}_val_auc', float('nan')):.4f}"
|
f" val_auc={result.get(f'{eval_stage}_val_auc', float('nan')):.4f}"
|
||||||
@@ -273,6 +300,7 @@ def main():
|
|||||||
"elapsed_s": round(time.time() - t0, 1),
|
"elapsed_s": round(time.time() - t0, 1),
|
||||||
"fold_results": fold_results,
|
"fold_results": fold_results,
|
||||||
}
|
}
|
||||||
|
out_dir.mkdir(parents=True, exist_ok=True)
|
||||||
summary_path = out_dir / "summary.json"
|
summary_path = out_dir / "summary.json"
|
||||||
summary_path.write_text(json.dumps(summary, indent=2))
|
summary_path.write_text(json.dumps(summary, indent=2))
|
||||||
print(f"\n{'='*60}", flush=True)
|
print(f"\n{'='*60}", flush=True)
|
||||||
@@ -280,6 +308,81 @@ def main():
|
|||||||
print(f"Test AUC: {summary['mean_test_auc']:.4f} ± {summary['std_test_auc']:.4f}", flush=True)
|
print(f"Test AUC: {summary['mean_test_auc']:.4f} ± {summary['std_test_auc']:.4f}", flush=True)
|
||||||
print(f"Saved: {summary_path}", flush=True)
|
print(f"Saved: {summary_path}", flush=True)
|
||||||
|
|
||||||
|
if save_predictions and eval_stage_preds:
|
||||||
|
# Collect all unique entity_ids across val+test sets of all folds.
|
||||||
|
seen, all_ids, id_to_y = set(), [], {}
|
||||||
|
for fp in eval_stage_preds:
|
||||||
|
for eid, y in zip(fp["val_ids"], fp["val_y"]):
|
||||||
|
k = str(eid)
|
||||||
|
if k not in seen:
|
||||||
|
seen.add(k); all_ids.append(eid)
|
||||||
|
id_to_y[k] = int(y)
|
||||||
|
if fp.get("test_ids"):
|
||||||
|
for eid, y in zip(fp["test_ids"], fp["test_y"]):
|
||||||
|
k = str(eid)
|
||||||
|
if k not in seen:
|
||||||
|
seen.add(k); all_ids.append(eid)
|
||||||
|
id_to_y[k] = int(y)
|
||||||
|
|
||||||
|
y_true = np.array([id_to_y.get(str(e), -1) for e in all_ids], dtype=np.int64)
|
||||||
|
store = PredictionStore(n_folds=len(eval_stage_preds), n_classes=num_classes)
|
||||||
|
store.register_phase(
|
||||||
|
phase=eval_stage,
|
||||||
|
entity_ids=all_ids,
|
||||||
|
y_true=y_true,
|
||||||
|
head_names=[f"{eval_stage}_head"],
|
||||||
|
n_epochs=1,
|
||||||
|
)
|
||||||
|
for fold_idx, fp in enumerate(eval_stage_preds):
|
||||||
|
store.record(eval_stage, fold_idx, 0, fp["val_ids"],
|
||||||
|
f"{eval_stage}_head", fp["val_p"])
|
||||||
|
store.set_split(eval_stage, fold_idx, fp["val_ids"], "val")
|
||||||
|
if fp.get("test_ids"):
|
||||||
|
store.record(eval_stage, fold_idx, 0, fp["test_ids"],
|
||||||
|
f"{eval_stage}_head", fp["test_p"])
|
||||||
|
store.set_split(eval_stage, fold_idx, fp["test_ids"], "test")
|
||||||
|
|
||||||
|
pred_path = out_dir / "predictions.h5"
|
||||||
|
store.save(pred_path)
|
||||||
|
print(f"Predictions saved: {pred_path}", flush=True)
|
||||||
|
|
||||||
|
if save_features and eval_stage_preds:
|
||||||
|
emb_dim = eval_stage_preds[0]["val_z"].shape[-1]
|
||||||
|
fstore = FeatureStore(n_folds=len(eval_stage_preds))
|
||||||
|
|
||||||
|
# Build entity_id / y_true universe (same as predictions).
|
||||||
|
seen, all_ids, id_to_y = set(), [], {}
|
||||||
|
for fp in eval_stage_preds:
|
||||||
|
for eid, y in zip(fp["val_ids"], fp["val_y"]):
|
||||||
|
k = str(eid)
|
||||||
|
if k not in seen:
|
||||||
|
seen.add(k); all_ids.append(eid)
|
||||||
|
id_to_y[k] = int(y)
|
||||||
|
if fp.get("test_ids"):
|
||||||
|
for eid, y in zip(fp["test_ids"], fp["test_y"]):
|
||||||
|
k = str(eid)
|
||||||
|
if k not in seen:
|
||||||
|
seen.add(k); all_ids.append(eid)
|
||||||
|
id_to_y[k] = int(y)
|
||||||
|
|
||||||
|
y_true = np.array([id_to_y.get(str(e), -1) for e in all_ids], dtype=np.int64)
|
||||||
|
fstore.register_phase(phase=eval_stage, entity_ids=all_ids, y_true=y_true)
|
||||||
|
fstore.register_head(phase=eval_stage, head=f"{eval_stage}_embedding",
|
||||||
|
n_epochs=1, embedding_dim=emb_dim)
|
||||||
|
|
||||||
|
for fold_idx, fp in enumerate(eval_stage_preds):
|
||||||
|
fstore.record(eval_stage, fold_idx, 0, fp["val_ids"],
|
||||||
|
f"{eval_stage}_embedding", fp["val_z"])
|
||||||
|
fstore.set_split(eval_stage, fold_idx, fp["val_ids"], "val")
|
||||||
|
if fp.get("test_ids") and fp.get("test_z") is not None:
|
||||||
|
fstore.record(eval_stage, fold_idx, 0, fp["test_ids"],
|
||||||
|
f"{eval_stage}_embedding", fp["test_z"])
|
||||||
|
fstore.set_split(eval_stage, fold_idx, fp["test_ids"], "test")
|
||||||
|
|
||||||
|
feat_path = out_dir / "features.h5"
|
||||||
|
fstore.save(feat_path)
|
||||||
|
print(f"Features saved: {feat_path}", flush=True)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
main()
|
||||||
|
|||||||
@@ -1,23 +1,25 @@
|
|||||||
{
|
{
|
||||||
"_notes": [
|
"_notes": [
|
||||||
"V4 stage-pipeline: warm → fusion stages with parallel head stages.",
|
"Shared-weight bilateral architecture: one img tower and one cd tower,",
|
||||||
"Bridges are pure embedding producers; heads are separate swappable stages.",
|
"each called for both eyes. nt trains on all eye-level data.",
|
||||||
"BCD-eligible heads are sampled during tower_warmup and main phases.",
|
"hb fuses nt(OD) and nt(OS) at patient level using the same frozen nt."
|
||||||
"eval_stage names the fusion stage whose primary head is used for final metrics."
|
|
||||||
],
|
],
|
||||||
|
|
||||||
"run_name": "v4/ensemble_fused",
|
"run_name": "v4/ensemble_fused",
|
||||||
"num_classes": 2,
|
"num_classes": 2,
|
||||||
"label_filter": [0, 1],
|
"label_filter": [
|
||||||
|
0,
|
||||||
|
1
|
||||||
|
],
|
||||||
"split_identity_level": 1,
|
"split_identity_level": 1,
|
||||||
"eval_stage": "hb",
|
"eval_stage": "hb",
|
||||||
"save_predictions": false,
|
"save_predictions": true,
|
||||||
"seed": 1234,
|
"seed": 1234,
|
||||||
"folds": 5,
|
"folds": 5,
|
||||||
"fold_seed": 100,
|
"fold_seed": 100,
|
||||||
"output_root": "v4/results",
|
"output_root": "v4/results",
|
||||||
"out_dir_tags": ["binary", "ntower"],
|
"out_dir_tags": [
|
||||||
|
"binary"
|
||||||
|
],
|
||||||
"data": {
|
"data": {
|
||||||
"module": "v4.classes.profiles.v4papila",
|
"module": "v4.classes.profiles.v4papila",
|
||||||
"args": {
|
"args": {
|
||||||
@@ -26,11 +28,12 @@
|
|||||||
"label_col": "Diagnosis",
|
"label_col": "Diagnosis",
|
||||||
"iop_corr_method": "ratio",
|
"iop_corr_method": "ratio",
|
||||||
"iop_drop_raw": true,
|
"iop_drop_raw": true,
|
||||||
"exclude_cols": ["Axial_Length"],
|
"exclude_cols": [
|
||||||
|
"Axial_Length"
|
||||||
|
],
|
||||||
"in_memory_cache": true
|
"in_memory_cache": true
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
||||||
"towers": [
|
"towers": [
|
||||||
{
|
{
|
||||||
"name": "img",
|
"name": "img",
|
||||||
@@ -53,12 +56,12 @@
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
|
|
||||||
"stages": [
|
"stages": [
|
||||||
{
|
{
|
||||||
"name": "cd_warm",
|
"name": "cd_warm",
|
||||||
"type": "warm",
|
"type": "warm",
|
||||||
"tower": "cd",
|
"tower": "cd",
|
||||||
|
"head_name": "cd_aux",
|
||||||
"level": "eye",
|
"level": "eye",
|
||||||
"epochs": 40
|
"epochs": 40
|
||||||
},
|
},
|
||||||
@@ -81,12 +84,20 @@
|
|||||||
"type": "fusion",
|
"type": "fusion",
|
||||||
"module": "v4.classes.bridges.fusion_bridge",
|
"module": "v4.classes.bridges.fusion_bridge",
|
||||||
"class": "FusionBridge",
|
"class": "FusionBridge",
|
||||||
"inputs": ["img", "cd"],
|
"inputs": [
|
||||||
|
"img",
|
||||||
|
"cd"
|
||||||
|
],
|
||||||
"level": "eye",
|
"level": "eye",
|
||||||
"epochs": 36,
|
"epochs": 36,
|
||||||
"train_towers": true,
|
"train_towers": true,
|
||||||
"warmup": { "tower_epochs": 3, "fused_epochs": 3 },
|
"warmup": {
|
||||||
"args": { "fusion_dim": 256 }
|
"tower_epochs": 3,
|
||||||
|
"fused_epochs": 3
|
||||||
|
},
|
||||||
|
"args": {
|
||||||
|
"fusion_dim": 256
|
||||||
|
}
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"name": "nt_head",
|
"name": "nt_head",
|
||||||
@@ -99,22 +110,30 @@
|
|||||||
"type": "fusion",
|
"type": "fusion",
|
||||||
"module": "v4.classes.bridges.hyperbridge",
|
"module": "v4.classes.bridges.hyperbridge",
|
||||||
"class": "HyperBridge",
|
"class": "HyperBridge",
|
||||||
"inputs": { "a": "nt", "b": "nt" },
|
"inputs": {
|
||||||
|
"a": "nt",
|
||||||
|
"b": "nt"
|
||||||
|
},
|
||||||
"level": "patient",
|
"level": "patient",
|
||||||
"epochs": 10,
|
"epochs": 10,
|
||||||
"args": { "hidden_dim": 256, "mode": "embedding_mlp" }
|
"args": {
|
||||||
|
"hidden_dim": 256,
|
||||||
|
"mode": "embedding_mlp"
|
||||||
|
}
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"name": "hb_head",
|
"name": "hb_head",
|
||||||
"type": "head",
|
"type": "head",
|
||||||
"input": "hb",
|
"input": "hb",
|
||||||
"train_with": "hb"
|
"train_with": "hb",
|
||||||
|
"args": {
|
||||||
|
"dropout": 0.3
|
||||||
|
}
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
|
|
||||||
"training": {
|
"training": {
|
||||||
"lr": 1e-4,
|
"lr": 1e-4,
|
||||||
"batch_size": 16,
|
"batch_size": 8,
|
||||||
"bcd_prob": 0.5,
|
"bcd_prob": 0.5,
|
||||||
"tune_binary_threshold": true
|
"tune_binary_threshold": true
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,149 @@
|
|||||||
|
{
|
||||||
|
"_notes": [
|
||||||
|
"Shared-weight bilateral architecture + GT geometry injection.",
|
||||||
|
"img supplies geometry_vectors via EPC; cd consumes them.",
|
||||||
|
"Same shared-weight structure as ensemble_fused.json."
|
||||||
|
],
|
||||||
|
"run_name": "v4/ensemble_fused_geom_gt",
|
||||||
|
"num_classes": 2,
|
||||||
|
"label_filter": [
|
||||||
|
0,
|
||||||
|
1
|
||||||
|
],
|
||||||
|
"split_identity_level": 1,
|
||||||
|
"eval_stage": "hb",
|
||||||
|
"save_predictions": true,
|
||||||
|
"seed": 1234,
|
||||||
|
"folds": 5,
|
||||||
|
"fold_seed": 100,
|
||||||
|
"output_root": "v4/results",
|
||||||
|
"out_dir_tags": [
|
||||||
|
"binary"
|
||||||
|
],
|
||||||
|
"data": {
|
||||||
|
"module": "v4.classes.profiles.v4papila",
|
||||||
|
"args": {
|
||||||
|
"image_dir": "Papila/FundusImages",
|
||||||
|
"clinical_dir": "Papila/ClinicalData",
|
||||||
|
"label_col": "Diagnosis",
|
||||||
|
"iop_corr_method": "ratio",
|
||||||
|
"iop_drop_raw": true,
|
||||||
|
"exclude_cols": [
|
||||||
|
"Axial_Length"
|
||||||
|
],
|
||||||
|
"in_memory_cache": true
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"towers": [
|
||||||
|
{
|
||||||
|
"name": "img",
|
||||||
|
"module": "v4.classes.towers.image_tower",
|
||||||
|
"class": "ImageEncoder",
|
||||||
|
"data_source": "image",
|
||||||
|
"epc_supplies": [
|
||||||
|
"geometry_vectors"
|
||||||
|
],
|
||||||
|
"args": {
|
||||||
|
"backbone": "refugelike",
|
||||||
|
"freeze_ratio": 0.0,
|
||||||
|
"augment": true,
|
||||||
|
"geometry_source": "gt",
|
||||||
|
"contour_dir": "Papila/ExpertsSegmentations/Contours"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "cd",
|
||||||
|
"module": "v4.classes.towers.clinical_tower",
|
||||||
|
"class": "ClinicalEncoder",
|
||||||
|
"data_source": "matrix",
|
||||||
|
"epc_requests": [
|
||||||
|
"geometry_vectors"
|
||||||
|
],
|
||||||
|
"args": {
|
||||||
|
"hidden_dim": 128,
|
||||||
|
"geom_dim": 5
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"stages": [
|
||||||
|
{
|
||||||
|
"name": "cd_warm",
|
||||||
|
"type": "warm",
|
||||||
|
"tower": "cd",
|
||||||
|
"head_name": "cd_aux",
|
||||||
|
"level": "eye",
|
||||||
|
"epochs": 40
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "img_aux",
|
||||||
|
"type": "head",
|
||||||
|
"input": "img",
|
||||||
|
"train_with": "nt",
|
||||||
|
"bcd": true
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "cd_aux",
|
||||||
|
"type": "head",
|
||||||
|
"input": "cd",
|
||||||
|
"train_with": "nt",
|
||||||
|
"bcd": true
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "nt",
|
||||||
|
"type": "fusion",
|
||||||
|
"module": "v4.classes.bridges.fusion_bridge",
|
||||||
|
"class": "FusionBridge",
|
||||||
|
"inputs": [
|
||||||
|
"img",
|
||||||
|
"cd"
|
||||||
|
],
|
||||||
|
"level": "eye",
|
||||||
|
"epochs": 36,
|
||||||
|
"train_towers": true,
|
||||||
|
"warmup": {
|
||||||
|
"tower_epochs": 3,
|
||||||
|
"fused_epochs": 3
|
||||||
|
},
|
||||||
|
"args": {
|
||||||
|
"fusion_dim": 256
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "nt_head",
|
||||||
|
"type": "head",
|
||||||
|
"input": "nt",
|
||||||
|
"train_with": "nt"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "hb",
|
||||||
|
"type": "fusion",
|
||||||
|
"module": "v4.classes.bridges.hyperbridge",
|
||||||
|
"class": "HyperBridge",
|
||||||
|
"inputs": {
|
||||||
|
"a": "nt",
|
||||||
|
"b": "nt"
|
||||||
|
},
|
||||||
|
"level": "patient",
|
||||||
|
"epochs": 10,
|
||||||
|
"args": {
|
||||||
|
"hidden_dim": 256,
|
||||||
|
"mode": "embedding_mlp"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "hb_head",
|
||||||
|
"type": "head",
|
||||||
|
"input": "hb",
|
||||||
|
"train_with": "hb",
|
||||||
|
"args": {
|
||||||
|
"dropout": 0.3
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"training": {
|
||||||
|
"lr": 1e-4,
|
||||||
|
"batch_size": 8,
|
||||||
|
"bcd_prob": 0.5,
|
||||||
|
"tune_binary_threshold": true
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
{
|
||||||
|
"_notes": [
|
||||||
|
"Standalone geometry tower: GeometrySegEncoder over UNet-derived seg maps,",
|
||||||
|
"MonoBridge passthrough, classification head. Same architecture as the",
|
||||||
|
"seg_cnn_unet_ft_mono test run, used as a 10-rep baseline for comparison",
|
||||||
|
"against tritower configurations."
|
||||||
|
],
|
||||||
|
"run_name": "v4/geometry_solo",
|
||||||
|
"num_classes": 2,
|
||||||
|
"label_filter": [0, 1],
|
||||||
|
"split_identity_level": 1,
|
||||||
|
"eval_stage": "geom_fuse",
|
||||||
|
"save_predictions": true,
|
||||||
|
"seed": 1234,
|
||||||
|
"folds": 5,
|
||||||
|
"fold_seed": 100,
|
||||||
|
"output_root": "v4/results",
|
||||||
|
"out_dir_tags": ["binary"],
|
||||||
|
|
||||||
|
"data": {
|
||||||
|
"module": "v4.classes.profiles.v4papila",
|
||||||
|
"args": {
|
||||||
|
"image_dir": "Papila/FundusImages",
|
||||||
|
"clinical_dir": "Papila/ClinicalData",
|
||||||
|
"label_col": "Diagnosis",
|
||||||
|
"iop_corr_method": "ratio",
|
||||||
|
"iop_drop_raw": true,
|
||||||
|
"exclude_cols": ["Axial_Length"],
|
||||||
|
"in_memory_cache": false
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
"towers": [
|
||||||
|
{
|
||||||
|
"name": "geom",
|
||||||
|
"module": "v4.classes.towers.geometry_tower",
|
||||||
|
"class": "GeometrySegEncoder",
|
||||||
|
"data_source": "image",
|
||||||
|
"args": {
|
||||||
|
"backbone": "resnet18",
|
||||||
|
"channels": 3,
|
||||||
|
"target_size": 224,
|
||||||
|
"augment": true,
|
||||||
|
"freeze_ratio": 0.0,
|
||||||
|
"seg_source": "unet",
|
||||||
|
"weights_path": "models/v2/refuge/segmentation/per_image/best.pt",
|
||||||
|
"contour_dir": "Papila/ExpertsSegmentations/Contours",
|
||||||
|
"unet_size": 512,
|
||||||
|
"normalize": "per_image",
|
||||||
|
"threshold": 0.5,
|
||||||
|
"crop_to_disc": true,
|
||||||
|
"finetune_epochs": 10,
|
||||||
|
"finetune_lr": 1e-5,
|
||||||
|
"finetune_batch_size": 4
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
|
||||||
|
"stages": [
|
||||||
|
{
|
||||||
|
"name": "geom_fuse",
|
||||||
|
"type": "fusion",
|
||||||
|
"module": "v4.classes.bridges.mono_bridge",
|
||||||
|
"class": "MonoBridge",
|
||||||
|
"inputs": ["geom"],
|
||||||
|
"level": "eye",
|
||||||
|
"epochs": 60,
|
||||||
|
"train_towers": true,
|
||||||
|
"args": { "use_ln": false }
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "geom_head",
|
||||||
|
"type": "head",
|
||||||
|
"input": "geom_fuse",
|
||||||
|
"train_with": "geom_fuse",
|
||||||
|
"args": { "dropout": 0.3 }
|
||||||
|
}
|
||||||
|
],
|
||||||
|
|
||||||
|
"training": {
|
||||||
|
"lr": 1e-4,
|
||||||
|
"batch_size": 16,
|
||||||
|
"tune_binary_threshold": true
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,163 @@
|
|||||||
|
{
|
||||||
|
"_notes": [
|
||||||
|
"Tri-tower shared-weight bilateral: img + cd + geom (UNet seg maps).",
|
||||||
|
"Default parameters; serves as both a 10-rep baseline and as the base",
|
||||||
|
"config for the bcd/cw/nt-epochs grid search.",
|
||||||
|
"geom is a GeometrySegEncoder over per-fold fine-tuned UNet predictions.",
|
||||||
|
"geom trains alongside img inside nt with the standard tower_warmup phase",
|
||||||
|
"(no separate geom_warm stage). hb fuses nt(OD) and nt(OS) at patient level."
|
||||||
|
],
|
||||||
|
"run_name": "v4/tritower",
|
||||||
|
"num_classes": 2,
|
||||||
|
"label_filter": [0, 1],
|
||||||
|
"split_identity_level": 1,
|
||||||
|
"eval_stage": "hb",
|
||||||
|
"save_predictions": true,
|
||||||
|
"seed": 1234,
|
||||||
|
"folds": 5,
|
||||||
|
"fold_seed": 100,
|
||||||
|
"output_root": "v4/results",
|
||||||
|
"out_dir_tags": ["binary"],
|
||||||
|
|
||||||
|
"data": {
|
||||||
|
"module": "v4.classes.profiles.v4papila",
|
||||||
|
"args": {
|
||||||
|
"image_dir": "Papila/FundusImages",
|
||||||
|
"clinical_dir": "Papila/ClinicalData",
|
||||||
|
"label_col": "Diagnosis",
|
||||||
|
"iop_corr_method": "ratio",
|
||||||
|
"iop_drop_raw": true,
|
||||||
|
"exclude_cols": ["Axial_Length"],
|
||||||
|
"in_memory_cache": true
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
"towers": [
|
||||||
|
{
|
||||||
|
"name": "img",
|
||||||
|
"module": "v4.classes.towers.image_tower",
|
||||||
|
"class": "ImageEncoder",
|
||||||
|
"data_source": "image",
|
||||||
|
"args": {
|
||||||
|
"backbone": "refugelike",
|
||||||
|
"freeze_ratio": 0.0,
|
||||||
|
"augment": true
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "cd",
|
||||||
|
"module": "v4.classes.towers.clinical_tower",
|
||||||
|
"class": "ClinicalEncoder",
|
||||||
|
"data_source": "matrix",
|
||||||
|
"args": {
|
||||||
|
"hidden_dim": 128
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "geom",
|
||||||
|
"module": "v4.classes.towers.geometry_tower",
|
||||||
|
"class": "GeometrySegEncoder",
|
||||||
|
"data_source": "image",
|
||||||
|
"args": {
|
||||||
|
"backbone": "resnet18",
|
||||||
|
"channels": 3,
|
||||||
|
"target_size": 224,
|
||||||
|
"augment": true,
|
||||||
|
"freeze_ratio": 0.0,
|
||||||
|
"seg_source": "unet",
|
||||||
|
"weights_path": "models/v2/refuge/segmentation/per_image/best.pt",
|
||||||
|
"contour_dir": "Papila/ExpertsSegmentations/Contours",
|
||||||
|
"unet_size": 512,
|
||||||
|
"normalize": "per_image",
|
||||||
|
"threshold": 0.5,
|
||||||
|
"crop_to_disc": true,
|
||||||
|
"finetune_epochs": 10,
|
||||||
|
"finetune_lr": 1e-5,
|
||||||
|
"finetune_batch_size": 4
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
|
||||||
|
"stages": [
|
||||||
|
{
|
||||||
|
"name": "cd_warm",
|
||||||
|
"type": "warm",
|
||||||
|
"tower": "cd",
|
||||||
|
"head_name": "cd_aux",
|
||||||
|
"level": "eye",
|
||||||
|
"epochs": 40
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "img_aux",
|
||||||
|
"type": "head",
|
||||||
|
"input": "img",
|
||||||
|
"train_with": "nt",
|
||||||
|
"bcd": true
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "cd_aux",
|
||||||
|
"type": "head",
|
||||||
|
"input": "cd",
|
||||||
|
"train_with": "nt",
|
||||||
|
"bcd": true
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "geom_aux",
|
||||||
|
"type": "head",
|
||||||
|
"input": "geom",
|
||||||
|
"train_with": "nt",
|
||||||
|
"bcd": true
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "nt",
|
||||||
|
"type": "fusion",
|
||||||
|
"module": "v4.classes.bridges.fusion_bridge",
|
||||||
|
"class": "FusionBridge",
|
||||||
|
"inputs": ["img", "cd", "geom"],
|
||||||
|
"level": "eye",
|
||||||
|
"epochs": 36,
|
||||||
|
"train_towers": true,
|
||||||
|
"warmup": {
|
||||||
|
"tower_epochs": 3,
|
||||||
|
"fused_epochs": 3
|
||||||
|
},
|
||||||
|
"args": {
|
||||||
|
"fusion_dim": 256
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "nt_head",
|
||||||
|
"type": "head",
|
||||||
|
"input": "nt",
|
||||||
|
"train_with": "nt"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "hb",
|
||||||
|
"type": "fusion",
|
||||||
|
"module": "v4.classes.bridges.hyperbridge",
|
||||||
|
"class": "HyperBridge",
|
||||||
|
"inputs": { "a": "nt", "b": "nt" },
|
||||||
|
"level": "patient",
|
||||||
|
"epochs": 10,
|
||||||
|
"args": {
|
||||||
|
"hidden_dim": 256,
|
||||||
|
"mode": "embedding_mlp"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "hb_head",
|
||||||
|
"type": "head",
|
||||||
|
"input": "hb",
|
||||||
|
"train_with": "hb",
|
||||||
|
"args": { "dropout": 0.3 }
|
||||||
|
}
|
||||||
|
],
|
||||||
|
|
||||||
|
"training": {
|
||||||
|
"lr": 1e-4,
|
||||||
|
"batch_size": 8,
|
||||||
|
"bcd_prob": 0.5,
|
||||||
|
"tune_binary_threshold": true,
|
||||||
|
"class_weighted": false
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,287 @@
|
|||||||
|
"""
|
||||||
|
v4 batch dispatch — merge a batch.json with a base config and submit jobs.
|
||||||
|
|
||||||
|
Each batch entry specifies only what differs from the base config. Each entry
|
||||||
|
is dispatched as N_REPS independent jobs (full 5-fold CV per rep, different seeds),
|
||||||
|
writing to {run_name}/rep00/, rep01/, etc.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m v4.distributed.batch_dispatch \\
|
||||||
|
--server http://apollo:8765 \\
|
||||||
|
--token <secret> \\
|
||||||
|
--config v4/configs/ensemble_fused.json \\
|
||||||
|
--batch v4/scripts/experiments/my_batch.json \\
|
||||||
|
[--reps 10] \\
|
||||||
|
[--seed-start 1234] \\
|
||||||
|
[--seed-step 100] \\
|
||||||
|
[--fold-seed-start 100] \\
|
||||||
|
[--fold-seed-step 100] \\
|
||||||
|
[--output-root v4/results] \\
|
||||||
|
[--priority 0] \\
|
||||||
|
[--dry-run]
|
||||||
|
|
||||||
|
batch.json format:
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"run_name": "experiments/lr_sweep/lr1e3", // required
|
||||||
|
"overrides": { // optional — deep-merged into base
|
||||||
|
"training": { "lr": 0.001 }
|
||||||
|
},
|
||||||
|
"stage_overrides": { // optional — patched by stage name
|
||||||
|
"nt": { "epochs": 40 }
|
||||||
|
},
|
||||||
|
"reps": 10, // optional — overrides --reps
|
||||||
|
"priority": 0 // optional
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
Global flags (can also be set via env vars):
|
||||||
|
--server HT_SERVER
|
||||||
|
--token HT_TOKEN
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import copy
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Config helpers
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def deep_merge(base: dict, override: dict) -> dict:
|
||||||
|
"""Recursively merge override into a copy of base.
|
||||||
|
|
||||||
|
- Dicts are merged recursively.
|
||||||
|
- All other types (scalars, lists) are replaced by the override value.
|
||||||
|
"""
|
||||||
|
result = copy.deepcopy(base)
|
||||||
|
for k, v in override.items():
|
||||||
|
if k in result and isinstance(result[k], dict) and isinstance(v, dict):
|
||||||
|
result[k] = deep_merge(result[k], v)
|
||||||
|
else:
|
||||||
|
result[k] = copy.deepcopy(v)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def apply_stage_overrides(stages: list[dict], stage_overrides: dict) -> list[dict]:
|
||||||
|
"""Patch individual stages by name without replacing the entire list."""
|
||||||
|
stages = copy.deepcopy(stages)
|
||||||
|
for stage in stages:
|
||||||
|
name = stage.get("name")
|
||||||
|
if name in stage_overrides:
|
||||||
|
merged = deep_merge(stage, stage_overrides[name])
|
||||||
|
stage.clear()
|
||||||
|
stage.update(merged)
|
||||||
|
return stages
|
||||||
|
|
||||||
|
|
||||||
|
def build_config(base_cfg: dict, entry: dict, rep: int, seed: int, fold_seed: int,
|
||||||
|
output_root: str) -> dict:
|
||||||
|
"""Produce the final merged config for one rep of one batch entry."""
|
||||||
|
cfg = copy.deepcopy(base_cfg)
|
||||||
|
|
||||||
|
# Deep-merge top-level overrides
|
||||||
|
cfg = deep_merge(cfg, entry.get("overrides", {}))
|
||||||
|
|
||||||
|
# Patch individual stages by name
|
||||||
|
if "stage_overrides" in entry and "stages" in cfg:
|
||||||
|
cfg["stages"] = apply_stage_overrides(cfg["stages"], entry["stage_overrides"])
|
||||||
|
|
||||||
|
# Stamp run_name, model seed, split seed, output_root.
|
||||||
|
base_run_name = entry["run_name"]
|
||||||
|
cfg["run_name"] = f"{base_run_name}/rep{rep:02d}"
|
||||||
|
cfg["seed"] = seed
|
||||||
|
cfg["fold_seed"] = fold_seed
|
||||||
|
cfg["output_root"] = output_root
|
||||||
|
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
|
def save_dispatched_config(cfg: dict, base_run_name: str, rep: int,
|
||||||
|
repo_root: Path) -> Path:
|
||||||
|
"""Write the merged config to v4/configs/dispatched/ and return its path."""
|
||||||
|
out_dir = repo_root / "v4" / "configs" / "dispatched" / base_run_name
|
||||||
|
out_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
path = out_dir / f"rep{rep:02d}.json"
|
||||||
|
path.write_text(json.dumps(cfg, indent=2))
|
||||||
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Server HTTP helper
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
class _API:
|
||||||
|
def __init__(self, base_url: str, token: str):
|
||||||
|
self.base_url = base_url.rstrip("/")
|
||||||
|
self._h = {"x-token": token}
|
||||||
|
|
||||||
|
def post(self, path: str, body: dict) -> dict:
|
||||||
|
r = requests.post(f"{self.base_url}{path}", headers=self._h,
|
||||||
|
json=body, timeout=10)
|
||||||
|
r.raise_for_status()
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Dispatch
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def dispatch_batch(
|
||||||
|
api: _API,
|
||||||
|
base_cfg: dict,
|
||||||
|
batch: list[dict],
|
||||||
|
*,
|
||||||
|
default_reps: int,
|
||||||
|
seed_start: int,
|
||||||
|
seed_step: int,
|
||||||
|
fold_seed_start: int,
|
||||||
|
fold_seed_step: int,
|
||||||
|
output_root: str,
|
||||||
|
default_priority: int,
|
||||||
|
server_path: str,
|
||||||
|
repo_root: Path,
|
||||||
|
dry_run: bool,
|
||||||
|
) -> None:
|
||||||
|
total = skipped = 0
|
||||||
|
for entry in batch:
|
||||||
|
run_name = entry["run_name"]
|
||||||
|
reps = entry.get("reps", default_reps)
|
||||||
|
priority = entry.get("priority", default_priority)
|
||||||
|
|
||||||
|
print(f"\n[dispatch] {run_name} ({reps} reps)")
|
||||||
|
|
||||||
|
for rep in range(reps):
|
||||||
|
seed = seed_start + rep * seed_step
|
||||||
|
fold_seed = fold_seed_start + rep * fold_seed_step
|
||||||
|
cfg = build_config(base_cfg, entry, rep, seed, fold_seed, output_root)
|
||||||
|
cfg_path = save_dispatched_config(cfg, run_name, rep, repo_root)
|
||||||
|
|
||||||
|
rel_path = cfg_path.relative_to(repo_root)
|
||||||
|
server_cfg_path = str(rel_path)
|
||||||
|
|
||||||
|
job_body = {
|
||||||
|
"run_name": run_name,
|
||||||
|
"module": "v4.classes.v4_hypertower",
|
||||||
|
"args": ["--config", server_cfg_path],
|
||||||
|
"output_dir": output_root,
|
||||||
|
"priority": priority,
|
||||||
|
}
|
||||||
|
|
||||||
|
if dry_run:
|
||||||
|
print(
|
||||||
|
f" [dry-run] rep{rep:02d} seed={seed} "
|
||||||
|
f"fold_seed={fold_seed} config={server_cfg_path}"
|
||||||
|
)
|
||||||
|
total += 1
|
||||||
|
else:
|
||||||
|
resp = api.post("/jobs", job_body)
|
||||||
|
if resp.get("skipped"):
|
||||||
|
print(f" rep{rep:02d} [skip — results exist on disk]")
|
||||||
|
skipped += 1
|
||||||
|
elif resp.get("duplicate"):
|
||||||
|
print(f" rep{rep:02d} [skip — already queued] job_id={resp['job_id']}")
|
||||||
|
skipped += 1
|
||||||
|
else:
|
||||||
|
print(
|
||||||
|
f" rep{rep:02d} seed={seed} fold_seed={fold_seed} "
|
||||||
|
f"job_id={resp['job_id']} config={server_cfg_path}"
|
||||||
|
)
|
||||||
|
total += 1
|
||||||
|
|
||||||
|
action = "would submit" if dry_run else "submitted"
|
||||||
|
skip_note = f" ({skipped} already queued/done, skipped)" if skipped else ""
|
||||||
|
print(f"\n[dispatch] {action} {total} jobs across {len(batch)} experiment(s){skip_note}")
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Entry point
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def main():
|
||||||
|
repo_root = Path(__file__).resolve().parents[2]
|
||||||
|
|
||||||
|
ap = argparse.ArgumentParser(
|
||||||
|
description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
|
||||||
|
)
|
||||||
|
ap.add_argument("--server", default=os.environ.get("HT_SERVER", ""),
|
||||||
|
help="Server URL (or set HT_SERVER)")
|
||||||
|
ap.add_argument("--token", default=os.environ.get("HT_TOKEN", ""),
|
||||||
|
help="Shared secret (or set HT_TOKEN)")
|
||||||
|
ap.add_argument("--config", required=True,
|
||||||
|
help="Base config JSON file")
|
||||||
|
ap.add_argument("--batch", required=True,
|
||||||
|
help="Batch JSON file listing experiments")
|
||||||
|
ap.add_argument("--reps", type=int, default=10,
|
||||||
|
help="Repetitions per experiment (default: 10)")
|
||||||
|
ap.add_argument("--seed-start", type=int, default=1234,
|
||||||
|
help="Seed for rep00 (default: 1234)")
|
||||||
|
ap.add_argument("--seed-step", type=int, default=100,
|
||||||
|
help="Seed increment per rep (default: 100)")
|
||||||
|
ap.add_argument("--fold-seed-start", type=int, default=100,
|
||||||
|
help="Split fold_seed for rep00 (default: 100; matches v3)")
|
||||||
|
ap.add_argument("--fold-seed-step", type=int, default=100,
|
||||||
|
help="Split fold_seed increment per rep (default: 100; matches v3)")
|
||||||
|
ap.add_argument("--output-root", default="v4/results",
|
||||||
|
help="Output root written into each config (default: v4/results)")
|
||||||
|
ap.add_argument("--server-path", default="",
|
||||||
|
help="Absolute path to hypertower root on server "
|
||||||
|
"(used to build config paths in job args; "
|
||||||
|
"if omitted, relative paths are used)")
|
||||||
|
ap.add_argument("--priority", type=int, default=0,
|
||||||
|
help="Default job priority (default: 0)")
|
||||||
|
ap.add_argument("--dry-run", action="store_true",
|
||||||
|
help="Print jobs without submitting")
|
||||||
|
args = ap.parse_args()
|
||||||
|
|
||||||
|
if not args.dry_run:
|
||||||
|
if not args.server:
|
||||||
|
ap.error("--server is required (or set HT_SERVER)")
|
||||||
|
if not args.token:
|
||||||
|
ap.error("--token is required (or set HT_TOKEN)")
|
||||||
|
|
||||||
|
cfg_path = Path(args.config)
|
||||||
|
if not cfg_path.is_absolute():
|
||||||
|
cfg_path = repo_root / cfg_path
|
||||||
|
base_cfg = json.loads(cfg_path.read_text())
|
||||||
|
|
||||||
|
batch_path = Path(args.batch)
|
||||||
|
if not batch_path.is_absolute():
|
||||||
|
batch_path = repo_root / batch_path
|
||||||
|
batch = json.loads(batch_path.read_text())
|
||||||
|
|
||||||
|
if not isinstance(batch, list):
|
||||||
|
sys.exit("batch.json must be a JSON array")
|
||||||
|
for i, entry in enumerate(batch):
|
||||||
|
if "run_name" not in entry:
|
||||||
|
sys.exit(f"batch entry {i} is missing required 'run_name'")
|
||||||
|
|
||||||
|
api = _API(args.server, args.token) if not args.dry_run else None
|
||||||
|
|
||||||
|
dispatch_batch(
|
||||||
|
api,
|
||||||
|
base_cfg,
|
||||||
|
batch,
|
||||||
|
default_reps=args.reps,
|
||||||
|
seed_start=args.seed_start,
|
||||||
|
seed_step=args.seed_step,
|
||||||
|
fold_seed_start=args.fold_seed_start,
|
||||||
|
fold_seed_step=args.fold_seed_step,
|
||||||
|
output_root=args.output_root,
|
||||||
|
default_priority=args.priority,
|
||||||
|
server_path=args.server_path,
|
||||||
|
repo_root=repo_root,
|
||||||
|
dry_run=args.dry_run,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,347 @@
|
|||||||
|
"""
|
||||||
|
HyperTower distributed job CLI — submit jobs, view status.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
# View all connected clients
|
||||||
|
python -m v4.distributed.cli clients
|
||||||
|
|
||||||
|
# View a specific client
|
||||||
|
python -m v4.distributed.cli clients <client_id>
|
||||||
|
|
||||||
|
# Live monitoring
|
||||||
|
python -m v4.distributed.cli clients --watch
|
||||||
|
|
||||||
|
# View jobs (optionally filter by state)
|
||||||
|
python -m v4.distributed.cli jobs [--state pending|running|done|failed]
|
||||||
|
|
||||||
|
# Submit a job
|
||||||
|
python -m v4.distributed.cli submit \\
|
||||||
|
--run-name v4/ensemble_fused \\
|
||||||
|
-- --config v4/configs/ensemble_fused.json
|
||||||
|
|
||||||
|
# Submit all jobs from a batch file (JSON)
|
||||||
|
python -m v4.distributed.cli submit-batch jobs.json
|
||||||
|
|
||||||
|
# Cancel a pending job
|
||||||
|
python -m v4.distributed.cli cancel <job_id>
|
||||||
|
|
||||||
|
Global flags (can also be set via env vars):
|
||||||
|
--server HT_SERVER e.g. http://apollo:8765
|
||||||
|
--token HT_TOKEN
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# HTTP helpers
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
class _API:
|
||||||
|
def __init__(self, base_url: str, token: str):
|
||||||
|
self.base_url = base_url.rstrip("/")
|
||||||
|
self._h = {"x-token": token}
|
||||||
|
|
||||||
|
def get(self, path: str, **params) -> object:
|
||||||
|
r = requests.get(f"{self.base_url}{path}", headers=self._h,
|
||||||
|
params=params, timeout=10)
|
||||||
|
r.raise_for_status()
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
def post(self, path: str, body: dict) -> object:
|
||||||
|
r = requests.post(f"{self.base_url}{path}", headers=self._h,
|
||||||
|
json=body, timeout=10)
|
||||||
|
r.raise_for_status()
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
def delete(self, path: str) -> object:
|
||||||
|
r = requests.delete(f"{self.base_url}{path}", headers=self._h, timeout=10)
|
||||||
|
r.raise_for_status()
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Formatting helpers
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _ago(ts: Optional[str]) -> str:
|
||||||
|
if not ts:
|
||||||
|
return "-"
|
||||||
|
try:
|
||||||
|
dt = datetime.fromisoformat(ts)
|
||||||
|
delta = datetime.now(timezone.utc) - dt
|
||||||
|
secs = int(delta.total_seconds())
|
||||||
|
if secs < 60:
|
||||||
|
return f"{secs}s ago"
|
||||||
|
elif secs < 3600:
|
||||||
|
return f"{secs//60}m ago"
|
||||||
|
else:
|
||||||
|
return f"{secs//3600}h{(secs%3600)//60}m ago"
|
||||||
|
except Exception:
|
||||||
|
return ts
|
||||||
|
|
||||||
|
|
||||||
|
def _table(rows: list[list[str]], headers: list[str]):
|
||||||
|
widths = [max(len(str(r[i])) for r in ([headers] + rows)) for i in range(len(headers))]
|
||||||
|
sep = " "
|
||||||
|
def _row(r):
|
||||||
|
return sep.join(str(r[i]).ljust(widths[i]) for i in range(len(r)))
|
||||||
|
print(_row(headers))
|
||||||
|
print("-" * (sum(widths) + len(sep) * (len(widths) - 1)))
|
||||||
|
for r in rows:
|
||||||
|
print(_row(r))
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Subcommands
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _clients_table(api: _API) -> str:
|
||||||
|
clients = api.get("/clients")
|
||||||
|
if not clients:
|
||||||
|
return "No clients connected."
|
||||||
|
rows = []
|
||||||
|
for c in clients:
|
||||||
|
s = c["status"]
|
||||||
|
parts = []
|
||||||
|
if s.get("fold") is not None:
|
||||||
|
parts.append(f"fold{s['fold']}")
|
||||||
|
if s.get("stage"):
|
||||||
|
parts.append(s["stage"])
|
||||||
|
if s.get("epoch") is not None:
|
||||||
|
parts.append(f"ep{s['epoch']}/{s.get('total_epochs', '?')}")
|
||||||
|
if s.get("last_val_auc") is not None:
|
||||||
|
parts.append(f"auc={s['last_val_auc']:.4f}")
|
||||||
|
prog = " ".join(parts) if parts else "-"
|
||||||
|
rows.append([
|
||||||
|
c["client_id"],
|
||||||
|
c["hostname"],
|
||||||
|
c["gpu_info"][:30],
|
||||||
|
s["state"],
|
||||||
|
s.get("run_name") or "-",
|
||||||
|
prog,
|
||||||
|
_ago(c["last_seen"]),
|
||||||
|
])
|
||||||
|
headers = ["ID", "HOST", "GPU", "STATE", "RUN", "PROGRESS", "SEEN"]
|
||||||
|
widths = [max(len(str(r[i])) for r in ([headers] + rows)) for i in range(len(headers))]
|
||||||
|
sep = " "
|
||||||
|
lines = []
|
||||||
|
lines.append(sep.join(str(h).ljust(widths[i]) for i, h in enumerate(headers)))
|
||||||
|
lines.append("-" * (sum(widths) + len(sep) * (len(widths) - 1)))
|
||||||
|
for r in rows:
|
||||||
|
lines.append(sep.join(str(r[i]).ljust(widths[i]) for i in range(len(r))))
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_clients(api: _API, args):
|
||||||
|
if hasattr(args, "client_id") and args.client_id:
|
||||||
|
data = api.get(f"/clients/{args.client_id}")
|
||||||
|
s = data["status"]
|
||||||
|
print(f"client_id : {data['client_id']}")
|
||||||
|
print(f"hostname : {data['hostname']}")
|
||||||
|
print(f"gpu : {data['gpu_info']}")
|
||||||
|
print(f"last_seen : {_ago(data['last_seen'])}")
|
||||||
|
print(f"state : {s['state']}")
|
||||||
|
if s.get("job_id"):
|
||||||
|
print(f"job : {s['job_id']} ({s.get('run_name', '')})")
|
||||||
|
if s.get("fold") is not None:
|
||||||
|
print(f"progress : fold {s['fold']} stage {s.get('stage', '?')} "
|
||||||
|
f"ep {s.get('epoch', '?')}/{s.get('total_epochs', '?')} "
|
||||||
|
f"val_auc={s.get('last_val_auc', '?')}")
|
||||||
|
return
|
||||||
|
|
||||||
|
watch = getattr(args, "watch", False)
|
||||||
|
interval = getattr(args, "interval", 5)
|
||||||
|
|
||||||
|
if not watch:
|
||||||
|
print(_clients_table(api))
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
now = datetime.now().strftime("%H:%M:%S")
|
||||||
|
print(f"\033[H\033[2J", end="")
|
||||||
|
print(f"HyperTower clients [{now}] (Ctrl-C to exit)\n")
|
||||||
|
print(_clients_table(api))
|
||||||
|
time.sleep(interval)
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
print("\nStopped.")
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_jobs(api: _API, args):
|
||||||
|
params = {}
|
||||||
|
if hasattr(args, "state") and args.state:
|
||||||
|
params["state"] = args.state
|
||||||
|
jobs = api.get("/jobs", **params)
|
||||||
|
if not jobs:
|
||||||
|
print("No jobs.")
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
clients = api.get("/clients")
|
||||||
|
id_to_client = {c["client_id"]: c for c in clients}
|
||||||
|
except Exception:
|
||||||
|
id_to_client = {}
|
||||||
|
|
||||||
|
rows = []
|
||||||
|
for j in jobs:
|
||||||
|
attempts = j.get("attempts", 0)
|
||||||
|
client_id = j.get("assigned_to")
|
||||||
|
client = id_to_client.get(client_id) if client_id else None
|
||||||
|
|
||||||
|
client_label = (client["hostname"] if client else client_id) if client_id else "-"
|
||||||
|
|
||||||
|
progress = "-"
|
||||||
|
if client:
|
||||||
|
s = client.get("status", {})
|
||||||
|
parts = []
|
||||||
|
if s.get("fold") is not None:
|
||||||
|
parts.append(f"fold{s['fold']}")
|
||||||
|
if s.get("stage"):
|
||||||
|
parts.append(s["stage"])
|
||||||
|
if s.get("epoch") is not None:
|
||||||
|
parts.append(f"ep{s['epoch']}/{s.get('total_epochs', '?')}")
|
||||||
|
if parts:
|
||||||
|
progress = " ".join(parts)
|
||||||
|
|
||||||
|
rows.append([
|
||||||
|
j["job_id"][:12],
|
||||||
|
j["run_name"],
|
||||||
|
j["state"],
|
||||||
|
f"{attempts}" if attempts else "-",
|
||||||
|
client_label,
|
||||||
|
progress,
|
||||||
|
_ago(j["created_at"]),
|
||||||
|
_ago(j.get("started_at")),
|
||||||
|
_ago(j.get("completed_at")),
|
||||||
|
])
|
||||||
|
_table(rows, ["JOB_ID", "RUN_NAME", "STATE", "TRIES", "CLIENT", "PROGRESS", "CREATED", "STARTED", "DONE"])
|
||||||
|
pending = sum(1 for j in jobs if j["state"] == "pending")
|
||||||
|
running = sum(1 for j in jobs if j["state"] == "running")
|
||||||
|
done = sum(1 for j in jobs if j["state"] == "done")
|
||||||
|
failed = sum(1 for j in jobs if j["state"] == "failed")
|
||||||
|
print(f"\n {len(jobs)} total | {pending} pending {running} running {done} done {failed} failed")
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_submit(api: _API, args):
|
||||||
|
body = {
|
||||||
|
"run_name": args.run_name,
|
||||||
|
"module": args.module,
|
||||||
|
"args": args.run_args,
|
||||||
|
"output_dir": args.output_dir,
|
||||||
|
"priority": args.priority,
|
||||||
|
}
|
||||||
|
resp = api.post("/jobs", body)
|
||||||
|
print(f"Queued job {resp['job_id']} ({args.run_name})")
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_submit_batch(api: _API, args):
|
||||||
|
with open(args.batch_file) as f:
|
||||||
|
jobs = json.load(f)
|
||||||
|
for job in jobs:
|
||||||
|
resp = api.post("/jobs", job)
|
||||||
|
print(f"Queued {resp['job_id']} ({job['run_name']})")
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_cancel(api: _API, args):
|
||||||
|
resp = api.delete(f"/jobs/{args.job_id}")
|
||||||
|
print(f"Cancelled {args.job_id}" if resp.get("ok") else resp)
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_clear(api: _API, args):
|
||||||
|
body: dict = {}
|
||||||
|
if args.all:
|
||||||
|
body["all"] = True
|
||||||
|
elif args.run_name:
|
||||||
|
body["run_name"] = args.run_name
|
||||||
|
else:
|
||||||
|
body["states"] = args.states or ["done", "failed", "cancelled"]
|
||||||
|
resp = api.post("/jobs/clear", body)
|
||||||
|
print(f"Cleared {resp['cleared']} jobs.")
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Parser
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def main():
|
||||||
|
ap = argparse.ArgumentParser(description=__doc__,
|
||||||
|
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||||
|
ap.add_argument("--server", default=os.environ.get("HT_SERVER", ""),
|
||||||
|
help="Server URL (or set HT_SERVER)")
|
||||||
|
ap.add_argument("--token", default=os.environ.get("HT_TOKEN", ""),
|
||||||
|
help="Shared secret (or set HT_TOKEN)")
|
||||||
|
|
||||||
|
sub = ap.add_subparsers(dest="cmd", required=True)
|
||||||
|
|
||||||
|
# clients
|
||||||
|
p_cl = sub.add_parser("clients", help="List clients or inspect one")
|
||||||
|
p_cl.add_argument("client_id", nargs="?")
|
||||||
|
p_cl.add_argument("--watch", "-w", action="store_true",
|
||||||
|
help="Live monitoring mode — refresh every --interval seconds")
|
||||||
|
p_cl.add_argument("--interval", "-n", type=int, default=5,
|
||||||
|
help="Refresh interval in seconds for --watch (default: 5)")
|
||||||
|
|
||||||
|
# jobs
|
||||||
|
p_j = sub.add_parser("jobs", help="List jobs")
|
||||||
|
p_j.add_argument("--state", choices=["pending", "running", "done", "failed", "cancelled"])
|
||||||
|
|
||||||
|
# submit
|
||||||
|
p_s = sub.add_parser("submit", help="Submit a single job")
|
||||||
|
p_s.add_argument("--run-name", required=True)
|
||||||
|
p_s.add_argument("--module", default="v4.classes.v4_hypertower")
|
||||||
|
p_s.add_argument("--output-dir", default="v4/results")
|
||||||
|
p_s.add_argument("--priority", type=int, default=0)
|
||||||
|
p_s.add_argument("run_args", nargs=argparse.REMAINDER,
|
||||||
|
help="Args after '--' are forwarded to the module")
|
||||||
|
|
||||||
|
# submit-batch
|
||||||
|
p_b = sub.add_parser("submit-batch", help="Submit jobs from a JSON file")
|
||||||
|
p_b.add_argument("batch_file")
|
||||||
|
|
||||||
|
# cancel
|
||||||
|
p_c = sub.add_parser("cancel", help="Cancel a pending job")
|
||||||
|
p_c.add_argument("job_id")
|
||||||
|
|
||||||
|
# clear
|
||||||
|
p_cl2 = sub.add_parser("clear", help="Delete jobs by run-name, state, or everything")
|
||||||
|
p_cl2.add_argument("--run-name", default=None, help="Delete all jobs with this run-name")
|
||||||
|
p_cl2.add_argument("--states", nargs="+",
|
||||||
|
default=None,
|
||||||
|
choices=["done", "failed", "cancelled", "pending", "running"],
|
||||||
|
help="Delete jobs in these states (default: done+failed+cancelled)")
|
||||||
|
p_cl2.add_argument("--all", action="store_true", help="Delete ALL jobs")
|
||||||
|
|
||||||
|
args = ap.parse_args()
|
||||||
|
|
||||||
|
if not args.server:
|
||||||
|
ap.error("--server is required (or set HT_SERVER)")
|
||||||
|
if not args.token:
|
||||||
|
ap.error("--token is required (or set HT_TOKEN)")
|
||||||
|
|
||||||
|
if hasattr(args, "run_args") and args.run_args and args.run_args[0] == "--":
|
||||||
|
args.run_args = args.run_args[1:]
|
||||||
|
|
||||||
|
api = _API(args.server, args.token)
|
||||||
|
|
||||||
|
dispatch = {
|
||||||
|
"clients": cmd_clients,
|
||||||
|
"jobs": cmd_jobs,
|
||||||
|
"submit": cmd_submit,
|
||||||
|
"submit-batch": cmd_submit_batch,
|
||||||
|
"cancel": cmd_cancel,
|
||||||
|
"clear": cmd_clear,
|
||||||
|
}
|
||||||
|
dispatch[args.cmd](api, args)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,407 @@
|
|||||||
|
"""
|
||||||
|
HyperTower distributed job client daemon.
|
||||||
|
|
||||||
|
Registers with the server, polls for jobs, syncs code, runs training,
|
||||||
|
uploads results, and loops. Parses stdout to stream live status.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m v4.distributed.client \\
|
||||||
|
--server http://apollo:8765 \\
|
||||||
|
--token <secret> \\
|
||||||
|
--server-ssh rpotter@apollo \\
|
||||||
|
--server-path /home/rpotter/hypertower \\
|
||||||
|
[--local-path ~/hypertower] \\
|
||||||
|
[--poll-interval 15]
|
||||||
|
|
||||||
|
Compatibility test (verify GPU env, 1-fold dry-run):
|
||||||
|
python -m v4.distributed.client ... --test --config v4/configs/ensemble_fused.json
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import shutil
|
||||||
|
import signal
|
||||||
|
import socket
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from .protocol import (
|
||||||
|
JobResult,
|
||||||
|
JobSpec,
|
||||||
|
PollResponse,
|
||||||
|
RegisterRequest,
|
||||||
|
RegisterResponse,
|
||||||
|
StatusPush,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Server HTTP wrapper
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
class _Server:
|
||||||
|
def __init__(self, base_url: str, token: str):
|
||||||
|
self.base_url = base_url.rstrip("/")
|
||||||
|
self._h = {"x-token": token}
|
||||||
|
self.client_id: str = ""
|
||||||
|
|
||||||
|
def _post(self, path: str, **kw) -> dict:
|
||||||
|
r = requests.post(f"{self.base_url}{path}", headers=self._h, timeout=15, **kw)
|
||||||
|
r.raise_for_status()
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
def register(self, hostname: str, gpu_info: str) -> str:
|
||||||
|
data = self._post("/register",
|
||||||
|
json={"hostname": hostname, "gpu_info": gpu_info})
|
||||||
|
self.client_id = data["client_id"]
|
||||||
|
self.hostname = hostname
|
||||||
|
self.gpu_info = gpu_info
|
||||||
|
return self.client_id
|
||||||
|
|
||||||
|
def _reregister(self):
|
||||||
|
try:
|
||||||
|
self._post("/register",
|
||||||
|
json={"hostname": self.hostname, "gpu_info": self.gpu_info},
|
||||||
|
params={"reuse_id": self.client_id})
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def poll(self) -> Optional[JobSpec]:
|
||||||
|
data = self._post("/poll", params={"client_id": self.client_id})
|
||||||
|
if data.get("please_reregister"):
|
||||||
|
self._reregister()
|
||||||
|
return JobSpec(**data["job"]) if data.get("job") else None
|
||||||
|
|
||||||
|
def push_status(self, status: StatusPush):
|
||||||
|
try:
|
||||||
|
r = requests.post(
|
||||||
|
f"{self.base_url}/status/{self.client_id}",
|
||||||
|
json=status.model_dump(),
|
||||||
|
headers=self._h,
|
||||||
|
timeout=5,
|
||||||
|
)
|
||||||
|
if r.ok and r.json().get("please_reregister"):
|
||||||
|
self._reregister()
|
||||||
|
except Exception:
|
||||||
|
pass # don't crash job on status push failure
|
||||||
|
|
||||||
|
def complete(self, job_id: str, success: bool, error_msg: Optional[str] = None):
|
||||||
|
self._post("/complete",
|
||||||
|
json={"job_id": job_id, "success": success, "error_msg": error_msg})
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# GPU info
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _gpu_info() -> str:
|
||||||
|
# NVIDIA
|
||||||
|
try:
|
||||||
|
out = subprocess.check_output(
|
||||||
|
["nvidia-smi", "--query-gpu=name,memory.total", "--format=csv,noheader"],
|
||||||
|
text=True, stderr=subprocess.DEVNULL,
|
||||||
|
).strip()
|
||||||
|
if out:
|
||||||
|
return " | ".join(out.splitlines())
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
# AMD
|
||||||
|
try:
|
||||||
|
out = subprocess.check_output(
|
||||||
|
["rocm-smi", "--showproductname", "--csv"],
|
||||||
|
text=True, stderr=subprocess.DEVNULL,
|
||||||
|
).strip().splitlines()
|
||||||
|
names = [l for l in out if l and not l.startswith("device")]
|
||||||
|
if names:
|
||||||
|
return "AMD: " + " | ".join(names)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
# AMD fallback
|
||||||
|
try:
|
||||||
|
out = subprocess.check_output(
|
||||||
|
["rocminfo"],
|
||||||
|
text=True, stderr=subprocess.DEVNULL,
|
||||||
|
)
|
||||||
|
names = [l.split(":", 1)[1].strip() for l in out.splitlines()
|
||||||
|
if "Marketing Name:" in l]
|
||||||
|
if names:
|
||||||
|
return "AMD: " + " | ".join(names)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return "no-gpu"
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# rsync helpers
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _rsync(src: str, dst: str, delete: bool = False):
|
||||||
|
cmd = ["rsync", "-az", "--info=progress2"]
|
||||||
|
if delete:
|
||||||
|
cmd.append("--delete")
|
||||||
|
cmd += [src, dst]
|
||||||
|
subprocess.run(cmd, check=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _sync_code(server_ssh: str, server_path: str, local_path: str):
|
||||||
|
"""Pull v4/ source from server → local (overwrites local changes)."""
|
||||||
|
src = f"{server_ssh}:{server_path}/v4/"
|
||||||
|
dst = f"{local_path}/v4/"
|
||||||
|
Path(dst).mkdir(parents=True, exist_ok=True)
|
||||||
|
_rsync(src, dst, delete=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _upload_results(server_ssh: str, server_path: str, local_path: str,
|
||||||
|
run_name: str, output_dir: str):
|
||||||
|
src = f"{local_path}/{output_dir}/{run_name}/"
|
||||||
|
dst = f"{server_ssh}:{server_path}/{output_dir}/{run_name}/"
|
||||||
|
remote_parent = f"{server_path}/{output_dir}/{Path(run_name).parent}"
|
||||||
|
subprocess.run(["ssh", server_ssh, f"mkdir -p '{remote_parent}'"], check=True)
|
||||||
|
_rsync(src, dst)
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_local(local_path: str, run_name: str, output_dir: str):
|
||||||
|
target = Path(local_path) / output_dir / run_name
|
||||||
|
if target.exists():
|
||||||
|
shutil.rmtree(target)
|
||||||
|
print(f"[client] cleaned {target}", flush=True)
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Stdout parsers (match v4_hypertower.py / fusion.py print format)
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
# " fold2 [nt] ep004/036 [fused_warmup ] loss=0.4321 acc=0.876 val_auc=0.7654"
|
||||||
|
_EP_RE = re.compile(
|
||||||
|
r"fold(\d+)\s+\[([^\]]+)\]\s+ep(\d+)/(\d+).*?val_auc=([0-9.nan]+)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
# "── fold 2/5 train_groups=..."
|
||||||
|
_FOLD_RE = re.compile(r"fold\s+(\d+)/\d+", re.IGNORECASE)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_line(line: str) -> dict:
|
||||||
|
"""Return any structured fields found in a stdout line."""
|
||||||
|
out = {}
|
||||||
|
m = _FOLD_RE.search(line)
|
||||||
|
if m:
|
||||||
|
out["fold"] = int(m.group(1))
|
||||||
|
m = _EP_RE.search(line)
|
||||||
|
if m:
|
||||||
|
out["fold"] = int(m.group(1))
|
||||||
|
out["stage"] = m.group(2)
|
||||||
|
out["epoch"] = int(m.group(3))
|
||||||
|
out["total_epochs"] = int(m.group(4))
|
||||||
|
try:
|
||||||
|
out["last_val_auc"] = float(m.group(5))
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Core job runner
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _run_job(job: JobSpec, server: _Server,
|
||||||
|
server_ssh: str, server_path: str, local_path: str,
|
||||||
|
no_sync: bool = False,
|
||||||
|
extra_args: list[str] | None = None) -> bool:
|
||||||
|
extra_args = extra_args or []
|
||||||
|
# 1. Sync code
|
||||||
|
if no_sync:
|
||||||
|
print(f"[client] skipping sync (--no-sync)", flush=True)
|
||||||
|
else:
|
||||||
|
print(f"[client] syncing v4/ from server...", flush=True)
|
||||||
|
server.push_status(StatusPush(state="syncing", job_id=job.job_id, run_name=job.run_name))
|
||||||
|
_sync_code(server_ssh, server_path, local_path)
|
||||||
|
|
||||||
|
# 2. Launch training subprocess
|
||||||
|
cmd = [sys.executable, "-m", job.module] + job.args + extra_args
|
||||||
|
print(f"[client] running: {' '.join(cmd)}", flush=True)
|
||||||
|
server.push_status(StatusPush(state="running", job_id=job.job_id, run_name=job.run_name))
|
||||||
|
|
||||||
|
log_dir = Path(local_path) / "v4" / "distributed" / "logs"
|
||||||
|
log_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
log_file = log_dir / f"job_{job.job_id}.log"
|
||||||
|
ctx: dict = {}
|
||||||
|
|
||||||
|
def _tail(path: Path):
|
||||||
|
with open(path, "r") as f:
|
||||||
|
while True:
|
||||||
|
raw = f.readline()
|
||||||
|
if raw:
|
||||||
|
print(raw, end="", flush=True)
|
||||||
|
info = _parse_line(raw)
|
||||||
|
ctx.update(info)
|
||||||
|
if "epoch" in info:
|
||||||
|
server.push_status(StatusPush(
|
||||||
|
state="running",
|
||||||
|
job_id=job.job_id,
|
||||||
|
run_name=job.run_name,
|
||||||
|
fold=ctx.get("fold"),
|
||||||
|
stage=ctx.get("stage"),
|
||||||
|
epoch=ctx.get("epoch"),
|
||||||
|
total_epochs=ctx.get("total_epochs"),
|
||||||
|
last_val_auc=ctx.get("last_val_auc"),
|
||||||
|
))
|
||||||
|
elif proc.poll() is not None:
|
||||||
|
for raw in f:
|
||||||
|
print(raw, end="", flush=True)
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
time.sleep(0.05)
|
||||||
|
|
||||||
|
with open(log_file, "w") as logf:
|
||||||
|
proc = subprocess.Popen(
|
||||||
|
cmd,
|
||||||
|
stdout=logf,
|
||||||
|
stderr=logf,
|
||||||
|
cwd=local_path,
|
||||||
|
start_new_session=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
tailer = threading.Thread(target=_tail, args=(log_file,), daemon=True)
|
||||||
|
tailer.start()
|
||||||
|
proc.wait()
|
||||||
|
tailer.join(timeout=5)
|
||||||
|
log_file.unlink(missing_ok=True)
|
||||||
|
|
||||||
|
success = proc.returncode == 0
|
||||||
|
|
||||||
|
if not success:
|
||||||
|
try:
|
||||||
|
os.killpg(os.getpgid(proc.pid), signal.SIGKILL)
|
||||||
|
except ProcessLookupError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if not success:
|
||||||
|
print(f"[client] job FAILED (rc={proc.returncode})", flush=True)
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 3. Upload results
|
||||||
|
if no_sync:
|
||||||
|
print(f"[client] skipping upload (--no-sync, results already local)", flush=True)
|
||||||
|
else:
|
||||||
|
print(f"[client] uploading results...", flush=True)
|
||||||
|
server.push_status(StatusPush(state="uploading", job_id=job.job_id, run_name=job.run_name))
|
||||||
|
_upload_results(server_ssh, server_path, local_path, job.run_name, job.output_dir)
|
||||||
|
_clean_local(local_path, job.run_name, job.output_dir)
|
||||||
|
|
||||||
|
print(f"[client] job {job.job_id} complete.", flush=True)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Compatibility test
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _compat_test(server: _Server, server_ssh: str, server_path: str,
|
||||||
|
local_path: str, config: str, extra_args: list[str],
|
||||||
|
no_sync: bool = False):
|
||||||
|
"""Run 1 fold to verify the env works end-to-end."""
|
||||||
|
print("[client] === compatibility test ===", flush=True)
|
||||||
|
job = JobSpec(
|
||||||
|
job_id="compat-test",
|
||||||
|
run_name="_compat_test",
|
||||||
|
module="v4.classes.v4_hypertower",
|
||||||
|
args=["--config", config, "--device", "cpu"] + extra_args,
|
||||||
|
output_dir="v4/results",
|
||||||
|
)
|
||||||
|
ok = _run_job(job, server, server_ssh, server_path, local_path,
|
||||||
|
no_sync=no_sync)
|
||||||
|
_clean_local(local_path, "_compat_test", "v4/results")
|
||||||
|
if ok:
|
||||||
|
print("[client] compatibility test PASSED ✓", flush=True)
|
||||||
|
else:
|
||||||
|
print("[client] compatibility test FAILED ✗", flush=True)
|
||||||
|
return ok
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Main daemon
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def main():
|
||||||
|
ap = argparse.ArgumentParser(
|
||||||
|
description=__doc__,
|
||||||
|
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||||
|
)
|
||||||
|
ap.add_argument("--server", required=True,
|
||||||
|
help="Server URL, e.g. http://apollo:8765")
|
||||||
|
ap.add_argument("--token", default=os.environ.get("HT_TOKEN", ""),
|
||||||
|
help="Shared secret (or set HT_TOKEN env var)")
|
||||||
|
ap.add_argument("--server-ssh", required=True,
|
||||||
|
help="SSH target for rsync, e.g. rpotter@apollo")
|
||||||
|
ap.add_argument("--server-path", required=True,
|
||||||
|
help="Absolute path to hypertower root on server")
|
||||||
|
ap.add_argument("--local-path",
|
||||||
|
default=str(Path.home() / "hypertower"),
|
||||||
|
help="Absolute path to hypertower root on this machine")
|
||||||
|
ap.add_argument("--poll-interval", type=int, default=15,
|
||||||
|
help="Seconds to wait between polls when idle")
|
||||||
|
ap.add_argument("--extra-args", nargs=argparse.REMAINDER, default=[],
|
||||||
|
help="Extra args appended to every job on this client. "
|
||||||
|
"Use -- to separate: --extra-args -- --device cpu")
|
||||||
|
ap.add_argument("--no-sync", action="store_true",
|
||||||
|
help="Skip rsync of v4/ before each job (use when client IS the server)")
|
||||||
|
ap.add_argument("--test", action="store_true",
|
||||||
|
help="Run 1-fold compatibility test and exit")
|
||||||
|
ap.add_argument("--config", default="v4/configs/ensemble_fused.json",
|
||||||
|
help="Config path for --test mode")
|
||||||
|
args = ap.parse_args()
|
||||||
|
|
||||||
|
if not args.token:
|
||||||
|
ap.error("--token is required (or set HT_TOKEN)")
|
||||||
|
|
||||||
|
hostname = socket.gethostname()
|
||||||
|
gpu_info = _gpu_info()
|
||||||
|
server = _Server(args.server, args.token)
|
||||||
|
|
||||||
|
client_id = server.register(hostname, gpu_info)
|
||||||
|
print(f"[client] registered as {client_id} ({hostname} | {gpu_info})", flush=True)
|
||||||
|
|
||||||
|
if args.test:
|
||||||
|
extra = [a for a in args.extra_args if a != "--"]
|
||||||
|
sys.exit(0 if _compat_test(
|
||||||
|
server, args.server_ssh, args.server_path,
|
||||||
|
args.local_path, args.config, extra,
|
||||||
|
no_sync=args.no_sync,
|
||||||
|
) else 1)
|
||||||
|
|
||||||
|
print(f"[client] polling every {args.poll_interval}s...", flush=True)
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
job = server.poll()
|
||||||
|
if job is None:
|
||||||
|
server.push_status(StatusPush(state="idle"))
|
||||||
|
time.sleep(args.poll_interval)
|
||||||
|
continue
|
||||||
|
|
||||||
|
success = _run_job(
|
||||||
|
job, server,
|
||||||
|
args.server_ssh, args.server_path, args.local_path,
|
||||||
|
no_sync=args.no_sync,
|
||||||
|
extra_args=[a for a in args.extra_args if a != "--"],
|
||||||
|
)
|
||||||
|
server.complete(job.job_id, success,
|
||||||
|
error_msg=None if success else "non-zero exit code")
|
||||||
|
server.push_status(StatusPush(state="idle"))
|
||||||
|
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
print("\n[client] shutting down", flush=True)
|
||||||
|
break
|
||||||
|
except Exception as exc:
|
||||||
|
print(f"[client] error: {exc}", flush=True)
|
||||||
|
time.sleep(args.poll_interval)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
[Unit]
|
||||||
|
Description=HyperTower distributed job server
|
||||||
|
After=network-online.target
|
||||||
|
Wants=network-online.target
|
||||||
|
|
||||||
|
[Service]
|
||||||
|
Type=simple
|
||||||
|
WorkingDirectory=/home/rpotter/hypertower
|
||||||
|
Environment=HT_TOKEN=hypertower
|
||||||
|
ExecStart=/home/rpotter/miniconda3/envs/fundus_imaging/bin/python -m v4.distributed.server --host 0.0.0.0 --port 8765
|
||||||
|
Restart=on-failure
|
||||||
|
RestartSec=5
|
||||||
|
StandardOutput=journal
|
||||||
|
StandardError=journal
|
||||||
|
|
||||||
|
[Install]
|
||||||
|
WantedBy=default.target
|
||||||
Binary file not shown.
@@ -0,0 +1,61 @@
|
|||||||
|
"""Shared data models for server/client communication."""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
|
||||||
|
class RegisterRequest(BaseModel):
|
||||||
|
hostname: str
|
||||||
|
gpu_info: str
|
||||||
|
|
||||||
|
|
||||||
|
class RegisterResponse(BaseModel):
|
||||||
|
client_id: str
|
||||||
|
|
||||||
|
|
||||||
|
class StatusPush(BaseModel):
|
||||||
|
state: str # idle | syncing | running | uploading | error
|
||||||
|
job_id: Optional[str] = None
|
||||||
|
run_name: Optional[str] = None
|
||||||
|
fold: Optional[int] = None
|
||||||
|
stage: Optional[str] = None
|
||||||
|
epoch: Optional[int] = None
|
||||||
|
total_epochs: Optional[int] = None
|
||||||
|
last_val_auc: Optional[float] = None
|
||||||
|
error: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class ClientInfo(BaseModel):
|
||||||
|
client_id: str
|
||||||
|
hostname: str
|
||||||
|
gpu_info: str
|
||||||
|
status: StatusPush
|
||||||
|
last_seen: str
|
||||||
|
|
||||||
|
|
||||||
|
class JobSpec(BaseModel):
|
||||||
|
job_id: str
|
||||||
|
run_name: str
|
||||||
|
module: str # e.g. "v4.classes.v4_hypertower"
|
||||||
|
args: list[str]
|
||||||
|
output_dir: str = "v4/results"
|
||||||
|
|
||||||
|
|
||||||
|
class PollResponse(BaseModel):
|
||||||
|
job: Optional[JobSpec] = None
|
||||||
|
please_reregister: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
class JobResult(BaseModel):
|
||||||
|
job_id: str
|
||||||
|
success: bool
|
||||||
|
error_msg: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class JobSubmit(BaseModel):
|
||||||
|
run_name: str
|
||||||
|
module: str = "v4.classes.v4_hypertower"
|
||||||
|
args: list[str]
|
||||||
|
output_dir: str = "v4/results"
|
||||||
|
priority: int = 0
|
||||||
@@ -0,0 +1,476 @@
|
|||||||
|
"""
|
||||||
|
HyperTower distributed job server.
|
||||||
|
|
||||||
|
Manages a SQLite job queue and a registry of connected clients.
|
||||||
|
Clients poll for work, push status updates, and report completion.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m v4.distributed.server --port 8765 --token <secret>
|
||||||
|
|
||||||
|
Environment:
|
||||||
|
HT_TOKEN — fallback if --token is not passed
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import sqlite3
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from fastapi import Depends, FastAPI, Header, HTTPException
|
||||||
|
import uvicorn
|
||||||
|
|
||||||
|
from .protocol import (
|
||||||
|
ClientInfo,
|
||||||
|
JobResult,
|
||||||
|
JobSpec,
|
||||||
|
JobSubmit,
|
||||||
|
PollResponse,
|
||||||
|
RegisterRequest,
|
||||||
|
RegisterResponse,
|
||||||
|
StatusPush,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Global state
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
_TOKEN: str = ""
|
||||||
|
_DB_PATH: Path = Path("v4/distributed/jobs.db")
|
||||||
|
_REPO_ROOT: Path = Path.cwd()
|
||||||
|
_CLIENT_TTL: int = 120 # seconds before a client is considered gone
|
||||||
|
_MAX_ATTEMPTS: int = 3 # max times a job is retried before being left as failed
|
||||||
|
|
||||||
|
_clients: dict[str, ClientInfo] = {}
|
||||||
|
_clients_lock = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
def _reap_stale_clients():
|
||||||
|
"""Background thread: remove silent clients and re-queue their running jobs."""
|
||||||
|
while True:
|
||||||
|
time.sleep(30)
|
||||||
|
cutoff = datetime.now(timezone.utc).timestamp() - _CLIENT_TTL
|
||||||
|
|
||||||
|
with _clients_lock:
|
||||||
|
stale = [
|
||||||
|
cid
|
||||||
|
for cid, c in _clients.items()
|
||||||
|
if datetime.fromisoformat(c.last_seen).timestamp() < cutoff
|
||||||
|
]
|
||||||
|
for cid in stale:
|
||||||
|
print(
|
||||||
|
f"[server] reaped stale client {cid} ({_clients[cid].hostname})",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
del _clients[cid]
|
||||||
|
known_ids = set(_clients.keys())
|
||||||
|
|
||||||
|
with _db() as conn:
|
||||||
|
rows = conn.execute(
|
||||||
|
"SELECT job_id, assigned_to FROM jobs WHERE state='running'"
|
||||||
|
).fetchall()
|
||||||
|
for row in rows:
|
||||||
|
if row["assigned_to"] not in known_ids:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE jobs SET state='pending', assigned_to=NULL, started_at=NULL "
|
||||||
|
"WHERE job_id=?",
|
||||||
|
(row["job_id"],),
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[server] re-queued job {row['job_id']} "
|
||||||
|
f"(client {row['assigned_to']} unknown)",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Database helpers
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _db():
|
||||||
|
conn = sqlite3.connect(str(_DB_PATH))
|
||||||
|
conn.row_factory = sqlite3.Row
|
||||||
|
try:
|
||||||
|
yield conn
|
||||||
|
conn.commit()
|
||||||
|
finally:
|
||||||
|
conn.close()
|
||||||
|
|
||||||
|
|
||||||
|
def _init_db():
|
||||||
|
_DB_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
with _db() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
CREATE TABLE IF NOT EXISTS jobs (
|
||||||
|
job_id TEXT PRIMARY KEY,
|
||||||
|
run_name TEXT NOT NULL,
|
||||||
|
module TEXT NOT NULL,
|
||||||
|
args TEXT NOT NULL, -- JSON list
|
||||||
|
output_dir TEXT NOT NULL DEFAULT 'v4/results',
|
||||||
|
state TEXT NOT NULL DEFAULT 'pending',
|
||||||
|
priority INTEGER NOT NULL DEFAULT 0,
|
||||||
|
assigned_to TEXT,
|
||||||
|
created_at TEXT NOT NULL,
|
||||||
|
started_at TEXT,
|
||||||
|
completed_at TEXT,
|
||||||
|
error_msg TEXT,
|
||||||
|
attempts INTEGER NOT NULL DEFAULT 0
|
||||||
|
)
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
conn.execute("CREATE INDEX IF NOT EXISTS idx_jobs_state ON jobs(state)")
|
||||||
|
conn.execute("CREATE INDEX IF NOT EXISTS idx_jobs_priority ON jobs(priority)")
|
||||||
|
conn.execute("CREATE INDEX IF NOT EXISTS idx_jobs_run_name ON jobs(run_name)")
|
||||||
|
conn.execute("CREATE UNIQUE INDEX IF NOT EXISTS idx_jobs_args ON jobs(args)")
|
||||||
|
try:
|
||||||
|
conn.execute(
|
||||||
|
"ALTER TABLE jobs ADD COLUMN attempts INTEGER NOT NULL DEFAULT 0"
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass # column already exists
|
||||||
|
|
||||||
|
|
||||||
|
def _now() -> str:
|
||||||
|
return datetime.now(timezone.utc).isoformat()
|
||||||
|
|
||||||
|
|
||||||
|
def _ensure_client(client_id: str, hostname: str = "", gpu_info: str = "") -> bool:
|
||||||
|
"""Re-register a client that survived a server restart.
|
||||||
|
Returns True if the client was unknown (placeholder created)."""
|
||||||
|
if client_id not in _clients:
|
||||||
|
_clients[client_id] = ClientInfo(
|
||||||
|
client_id=client_id,
|
||||||
|
hostname=hostname or client_id,
|
||||||
|
gpu_info=gpu_info or "unknown",
|
||||||
|
status=StatusPush(state="idle"),
|
||||||
|
last_seen=_now(),
|
||||||
|
)
|
||||||
|
print(f"[server] re-registered {client_id} (survived restart)", flush=True)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# FastAPI app
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
app = FastAPI(title="HyperTower Job Server")
|
||||||
|
|
||||||
|
|
||||||
|
def _check_token(x_token: str = Header(...)):
|
||||||
|
if x_token != _TOKEN:
|
||||||
|
raise HTTPException(status_code=403, detail="Invalid token")
|
||||||
|
|
||||||
|
|
||||||
|
# ── Registration ──────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
@app.post(
|
||||||
|
"/register", response_model=RegisterResponse, dependencies=[Depends(_check_token)]
|
||||||
|
)
|
||||||
|
def register(req: RegisterRequest, reuse_id: Optional[str] = None):
|
||||||
|
with _clients_lock:
|
||||||
|
client_id = (
|
||||||
|
reuse_id if (reuse_id and reuse_id in _clients) else str(uuid.uuid4())[:8]
|
||||||
|
)
|
||||||
|
existing_status = (
|
||||||
|
_clients[client_id].status
|
||||||
|
if client_id in _clients
|
||||||
|
else StatusPush(state="idle")
|
||||||
|
)
|
||||||
|
_clients[client_id] = ClientInfo(
|
||||||
|
client_id=client_id,
|
||||||
|
hostname=req.hostname,
|
||||||
|
gpu_info=req.gpu_info,
|
||||||
|
status=existing_status,
|
||||||
|
last_seen=_now(),
|
||||||
|
)
|
||||||
|
action = "re-registered" if reuse_id else "registered"
|
||||||
|
print(
|
||||||
|
f"[server] {action} {client_id} ({req.hostname} | {req.gpu_info})", flush=True
|
||||||
|
)
|
||||||
|
return RegisterResponse(client_id=client_id)
|
||||||
|
|
||||||
|
|
||||||
|
# ── Job polling ───────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/poll", response_model=PollResponse, dependencies=[Depends(_check_token)])
|
||||||
|
def poll(client_id: str):
|
||||||
|
with _clients_lock:
|
||||||
|
needs_reregister = _ensure_client(client_id)
|
||||||
|
_clients[client_id].last_seen = _now()
|
||||||
|
|
||||||
|
with _db() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT * FROM jobs WHERE state='pending' "
|
||||||
|
"ORDER BY priority DESC, created_at ASC LIMIT 1"
|
||||||
|
).fetchone()
|
||||||
|
|
||||||
|
if row is None:
|
||||||
|
return PollResponse(job=None, please_reregister=needs_reregister)
|
||||||
|
|
||||||
|
job_id = row["job_id"]
|
||||||
|
cur = conn.execute(
|
||||||
|
"UPDATE jobs SET state='pending', assigned_to=NULL, started_at=NULL "
|
||||||
|
"WHERE assigned_to=? AND state='running' AND job_id!=?",
|
||||||
|
(client_id, job_id),
|
||||||
|
)
|
||||||
|
if cur.rowcount:
|
||||||
|
print(
|
||||||
|
f"[server] reset {cur.rowcount} orphaned running job(s) for {client_id}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE jobs SET state='running', assigned_to=?, started_at=? WHERE job_id=?",
|
||||||
|
(client_id, _now(), job_id),
|
||||||
|
)
|
||||||
|
|
||||||
|
job = JobSpec(
|
||||||
|
job_id=job_id,
|
||||||
|
run_name=row["run_name"],
|
||||||
|
module=row["module"],
|
||||||
|
args=json.loads(row["args"]),
|
||||||
|
output_dir=row["output_dir"],
|
||||||
|
)
|
||||||
|
|
||||||
|
with _clients_lock:
|
||||||
|
_clients[client_id].status = StatusPush(
|
||||||
|
state="syncing", job_id=job_id, run_name=row["run_name"]
|
||||||
|
)
|
||||||
|
|
||||||
|
print(f"[server] dispatched {job_id} ({row['run_name']}) → {client_id}", flush=True)
|
||||||
|
return PollResponse(job=job, please_reregister=needs_reregister)
|
||||||
|
|
||||||
|
|
||||||
|
# ── Status ────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/status/{client_id}", dependencies=[Depends(_check_token)])
|
||||||
|
def push_status(client_id: str, status: StatusPush):
|
||||||
|
with _clients_lock:
|
||||||
|
needs_reregister = _ensure_client(client_id)
|
||||||
|
_clients[client_id].status = status
|
||||||
|
_clients[client_id].last_seen = _now()
|
||||||
|
return {"ok": True, "please_reregister": needs_reregister}
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/clients", dependencies=[Depends(_check_token)])
|
||||||
|
def list_clients():
|
||||||
|
with _clients_lock:
|
||||||
|
return list(_clients.values())
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/clients/{client_id}", dependencies=[Depends(_check_token)])
|
||||||
|
def get_client(client_id: str):
|
||||||
|
with _clients_lock:
|
||||||
|
if client_id not in _clients:
|
||||||
|
raise HTTPException(status_code=404, detail="Unknown client")
|
||||||
|
return _clients[client_id]
|
||||||
|
|
||||||
|
|
||||||
|
# ── Job completion ────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/complete", dependencies=[Depends(_check_token)])
|
||||||
|
def complete(result: JobResult):
|
||||||
|
with _db() as conn:
|
||||||
|
if result.success:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE jobs SET state='done', completed_at=?, error_msg=NULL WHERE job_id=?",
|
||||||
|
(_now(), result.job_id),
|
||||||
|
)
|
||||||
|
print(f"[server] job {result.job_id} → done", flush=True)
|
||||||
|
|
||||||
|
run_row = conn.execute(
|
||||||
|
"SELECT run_name FROM jobs WHERE job_id=?", (result.job_id,)
|
||||||
|
).fetchone()
|
||||||
|
if run_row:
|
||||||
|
run_name = run_row["run_name"]
|
||||||
|
remaining = conn.execute(
|
||||||
|
"SELECT COUNT(*) FROM jobs WHERE run_name=? AND state != 'done'",
|
||||||
|
(run_name,),
|
||||||
|
).fetchone()[0]
|
||||||
|
if remaining == 0:
|
||||||
|
total = conn.execute(
|
||||||
|
"SELECT COUNT(*) FROM jobs WHERE run_name=?", (run_name,)
|
||||||
|
).fetchone()[0]
|
||||||
|
conn.execute("DELETE FROM jobs WHERE run_name=?", (run_name,))
|
||||||
|
print(
|
||||||
|
f"[server] run '{run_name}' complete ({total} jobs) — cleared",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT attempts FROM jobs WHERE job_id=?", (result.job_id,)
|
||||||
|
).fetchone()
|
||||||
|
attempts = (row["attempts"] if row else 0) + 1
|
||||||
|
if attempts < _MAX_ATTEMPTS:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE jobs SET state='pending', assigned_to=NULL, started_at=NULL, "
|
||||||
|
"attempts=?, error_msg=? WHERE job_id=?",
|
||||||
|
(attempts, result.error_msg, result.job_id),
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[server] job {result.job_id} failed (attempt {attempts}/{_MAX_ATTEMPTS}), "
|
||||||
|
f"re-queuing",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE jobs SET state='failed', completed_at=?, attempts=?, error_msg=? "
|
||||||
|
"WHERE job_id=?",
|
||||||
|
(_now(), attempts, result.error_msg, result.job_id),
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[server] job {result.job_id} failed permanently after "
|
||||||
|
f"{attempts} attempts",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
return {"ok": True}
|
||||||
|
|
||||||
|
|
||||||
|
# ── Job queue management ──────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/jobs", dependencies=[Depends(_check_token)])
|
||||||
|
def submit_job(job: JobSubmit):
|
||||||
|
result_dir = _REPO_ROOT / job.output_dir / job.run_name
|
||||||
|
if result_dir.exists() and any(result_dir.rglob("summary.json")):
|
||||||
|
print(f"[server] skipped {job.run_name} (results exist on disk)", flush=True)
|
||||||
|
return {"job_id": "", "duplicate": False, "skipped": True}
|
||||||
|
|
||||||
|
args_json = json.dumps(job.args)
|
||||||
|
job_id = str(uuid.uuid4())[:12]
|
||||||
|
with _db() as conn:
|
||||||
|
cur = conn.execute(
|
||||||
|
"INSERT OR IGNORE INTO jobs "
|
||||||
|
"(job_id, run_name, module, args, output_dir, priority, created_at) "
|
||||||
|
"VALUES (?,?,?,?,?,?,?)",
|
||||||
|
(job_id, job.run_name, job.module, args_json,
|
||||||
|
job.output_dir, job.priority, _now()),
|
||||||
|
)
|
||||||
|
if cur.rowcount == 0:
|
||||||
|
existing = conn.execute(
|
||||||
|
"SELECT job_id FROM jobs WHERE args=?", (args_json,)
|
||||||
|
).fetchone()
|
||||||
|
job_id = existing["job_id"]
|
||||||
|
print(f"[server] duplicate ignored ({job.run_name}) → {job_id}", flush=True)
|
||||||
|
return {"job_id": job_id, "duplicate": True}
|
||||||
|
print(f"[server] queued {job_id} ({job.run_name})", flush=True)
|
||||||
|
return {"job_id": job_id, "duplicate": False}
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/jobs", dependencies=[Depends(_check_token)])
|
||||||
|
def list_jobs(state: Optional[str] = None):
|
||||||
|
with _db() as conn:
|
||||||
|
if state:
|
||||||
|
rows = conn.execute(
|
||||||
|
"SELECT * FROM jobs WHERE state=? ORDER BY created_at DESC", (state,)
|
||||||
|
).fetchall()
|
||||||
|
else:
|
||||||
|
rows = conn.execute(
|
||||||
|
"SELECT * FROM jobs ORDER BY created_at DESC"
|
||||||
|
).fetchall()
|
||||||
|
return [dict(r) for r in rows]
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/jobs/clear", dependencies=[Depends(_check_token)])
|
||||||
|
def clear_jobs(body: dict):
|
||||||
|
with _db() as conn:
|
||||||
|
if body.get("all"):
|
||||||
|
cur = conn.execute("DELETE FROM jobs")
|
||||||
|
elif body.get("run_name"):
|
||||||
|
cur = conn.execute("DELETE FROM jobs WHERE run_name=?", (body["run_name"],))
|
||||||
|
else:
|
||||||
|
states = body.get("states", ["done", "failed", "cancelled"])
|
||||||
|
placeholders = ",".join("?" * len(states))
|
||||||
|
cur = conn.execute(
|
||||||
|
f"DELETE FROM jobs WHERE state IN ({placeholders})", states
|
||||||
|
)
|
||||||
|
print(f"[server] cleared {cur.rowcount} jobs", flush=True)
|
||||||
|
return {"cleared": cur.rowcount}
|
||||||
|
|
||||||
|
|
||||||
|
@app.delete("/jobs/{job_id}", dependencies=[Depends(_check_token)])
|
||||||
|
def cancel_job(job_id: str):
|
||||||
|
with _db() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE jobs SET state='cancelled' WHERE job_id=? AND state='pending'",
|
||||||
|
(job_id,),
|
||||||
|
)
|
||||||
|
return {"ok": True}
|
||||||
|
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Entry point
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
ap = argparse.ArgumentParser(
|
||||||
|
description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
|
||||||
|
)
|
||||||
|
ap.add_argument("--port", type=int, default=8765)
|
||||||
|
ap.add_argument("--host", default="0.0.0.0")
|
||||||
|
ap.add_argument(
|
||||||
|
"--token",
|
||||||
|
default=os.environ.get("HT_TOKEN", ""),
|
||||||
|
help="Shared secret (or set HT_TOKEN env var)",
|
||||||
|
)
|
||||||
|
ap.add_argument(
|
||||||
|
"--db", default="v4/distributed/jobs.db", help="Path to SQLite job database"
|
||||||
|
)
|
||||||
|
ap.add_argument(
|
||||||
|
"--client-ttl",
|
||||||
|
type=int,
|
||||||
|
default=120,
|
||||||
|
help="Seconds of silence before a client is reaped (default: 120)",
|
||||||
|
)
|
||||||
|
ap.add_argument(
|
||||||
|
"--max-attempts",
|
||||||
|
type=int,
|
||||||
|
default=3,
|
||||||
|
help="Max times a failed job is retried before being left as failed (default: 3)",
|
||||||
|
)
|
||||||
|
ap.add_argument(
|
||||||
|
"--root",
|
||||||
|
default="",
|
||||||
|
help="Repo root for results-existence checks (default: cwd)",
|
||||||
|
)
|
||||||
|
args = ap.parse_args()
|
||||||
|
|
||||||
|
if not args.token:
|
||||||
|
ap.error("--token is required (or set HT_TOKEN)")
|
||||||
|
|
||||||
|
global _TOKEN, _DB_PATH, _REPO_ROOT, _CLIENT_TTL, _MAX_ATTEMPTS
|
||||||
|
_TOKEN = args.token
|
||||||
|
_DB_PATH = Path(args.db)
|
||||||
|
_REPO_ROOT = Path(args.root).resolve() if args.root else Path.cwd()
|
||||||
|
_CLIENT_TTL = args.client_ttl
|
||||||
|
_MAX_ATTEMPTS = args.max_attempts
|
||||||
|
_init_db()
|
||||||
|
|
||||||
|
reaper = threading.Thread(target=_reap_stale_clients, daemon=True)
|
||||||
|
reaper.start()
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"[server] listening on {args.host}:{args.port} client_ttl={_CLIENT_TTL}s",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
uvicorn.run(app, host=args.host, port=args.port, log_level="warning")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,147 @@
|
|||||||
|
# Distributed Server Cheat Sheet
|
||||||
|
|
||||||
|
All commands assume server is running on hades at port 8765.
|
||||||
|
|
||||||
|
## Start Server
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.server --token hypertower
|
||||||
|
```
|
||||||
|
Run inside tmux so it survives disconnects:
|
||||||
|
```bash
|
||||||
|
tmux new -s htserver
|
||||||
|
python -m v4.distributed.server --token hypertower
|
||||||
|
# Ctrl-B D to detach
|
||||||
|
tmux attach -t htserver # reattach later
|
||||||
|
```
|
||||||
|
|
||||||
|
## Local Workflow (hades as server + client)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Terminal 1 — server
|
||||||
|
python -m v4.distributed.server --token hypertower
|
||||||
|
|
||||||
|
# Terminal 2 — client (skip rsync, results already local)
|
||||||
|
python -m v4.distributed.client \
|
||||||
|
--server http://hades:8765 --token hypertower \
|
||||||
|
--server-ssh ignored --server-path ignored \
|
||||||
|
--local-path /home/rpotter/hypertower \
|
||||||
|
--no-sync
|
||||||
|
|
||||||
|
# Terminal 3 — dispatch (exits after queuing; client picks up jobs)
|
||||||
|
python -m v4.distributed.batch_dispatch \
|
||||||
|
--server http://hades:8765 --token hypertower \
|
||||||
|
--config v4/configs/ensemble_fused.json \
|
||||||
|
--batch v4/scripts/experiments/my_batch.json
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Start Clients
|
||||||
|
|
||||||
|
**Hades (server-local, no sync):**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.client \
|
||||||
|
--server http://hades:8765 --token hypertower \
|
||||||
|
--server-ssh ignored --server-path ignored \
|
||||||
|
--local-path /home/rpotter/hypertower \
|
||||||
|
--no-sync
|
||||||
|
```
|
||||||
|
|
||||||
|
**Apollo (remote client):**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.client \
|
||||||
|
--server http://hades:8765 --token hypertower \
|
||||||
|
--server-ssh rpotter@hades \
|
||||||
|
--server-path /home/rpotter/hypertower \
|
||||||
|
--local-path /home/odin/hypertower
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Monitoring
|
||||||
|
|
||||||
|
**Live client monitor (refreshes every 5s):**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower clients --watch
|
||||||
|
```
|
||||||
|
|
||||||
|
**Faster refresh:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower clients --watch --interval 2
|
||||||
|
```
|
||||||
|
|
||||||
|
**Inspect a single client:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower clients <client_id>
|
||||||
|
```
|
||||||
|
|
||||||
|
**View job queue:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower jobs
|
||||||
|
```
|
||||||
|
|
||||||
|
**Filter by state:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower jobs --state pending
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower jobs --state running
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower jobs --state failed
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Submitting Jobs
|
||||||
|
|
||||||
|
**Batch dispatch (dry run first):**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.batch_dispatch \
|
||||||
|
--server http://hades:8765 --token hypertower \
|
||||||
|
--config v4/configs/ensemble_fused.json \
|
||||||
|
--batch v4/scripts/experiments/fusion_dim_sweep.json \
|
||||||
|
--dry-run
|
||||||
|
```
|
||||||
|
|
||||||
|
**Live submit:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.batch_dispatch \
|
||||||
|
--server http://hades:8765 --token hypertower \
|
||||||
|
--config v4/configs/ensemble_fused.json \
|
||||||
|
--batch v4/scripts/experiments/fusion_dim_sweep.json
|
||||||
|
```
|
||||||
|
|
||||||
|
**Fewer reps (e.g. quick test):**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.batch_dispatch \
|
||||||
|
--server http://hades:8765 --token hypertower \
|
||||||
|
--config v4/configs/ensemble_fused.json \
|
||||||
|
--batch v4/scripts/experiments/fusion_dim_sweep.json \
|
||||||
|
--reps 3
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Queue Management
|
||||||
|
|
||||||
|
**Clear failed jobs:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower clear --states failed
|
||||||
|
```
|
||||||
|
|
||||||
|
**Clear running jobs (orphan cleanup):**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower clear --states running
|
||||||
|
```
|
||||||
|
|
||||||
|
**Clear all jobs:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower clear --all
|
||||||
|
```
|
||||||
|
|
||||||
|
**Clear by run name:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower clear --run-name experiments/fusion_dim_sweep/dim256
|
||||||
|
```
|
||||||
|
|
||||||
|
**Cancel a specific job:**
|
||||||
|
```bash
|
||||||
|
python -m v4.distributed.cli --server http://hades:8765 --token hypertower cancel <job_id>
|
||||||
|
```
|
||||||
@@ -1,146 +0,0 @@
|
|||||||
{
|
|
||||||
"run_name": "v4/ensemble_fused",
|
|
||||||
"config": {
|
|
||||||
"_notes": [
|
|
||||||
"V4 ensemble_fused: img + cd towers, HTFusion Stage 1, HyperBridge Stage 2.",
|
|
||||||
"Matches phase5/embedding_mlp_head setup for direct comparison.",
|
|
||||||
"data_source resolves against the PapilaBundle returned by build_data.",
|
|
||||||
"image_dir / clinical_dir are project-relative; orchestrator resolves against REPO_ROOT."
|
|
||||||
],
|
|
||||||
"run_name": "v4/ensemble_fused",
|
|
||||||
"eval_mode": "binary",
|
|
||||||
"split_identity_level": 1,
|
|
||||||
"epochs": 30,
|
|
||||||
"fusion_epochs": 10,
|
|
||||||
"folds": 5,
|
|
||||||
"fold_seed": 100,
|
|
||||||
"seed": 1234,
|
|
||||||
"output_root": "v4/results",
|
|
||||||
"data": {
|
|
||||||
"module": "v4.classes.profiles.v4papila",
|
|
||||||
"args": {
|
|
||||||
"image_dir": "Papila/FundusImages",
|
|
||||||
"clinical_dir": "Papila/ClinicalData",
|
|
||||||
"label_col": "Diagnosis",
|
|
||||||
"iop_corr_method": "ratio",
|
|
||||||
"iop_drop_raw": true,
|
|
||||||
"exclude_cols": [
|
|
||||||
"Axial_Length"
|
|
||||||
],
|
|
||||||
"in_memory_cache": true
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"towers": [
|
|
||||||
{
|
|
||||||
"name": "img",
|
|
||||||
"module": "v4.classes.towers.image_tower",
|
|
||||||
"class": "ImageEncoder",
|
|
||||||
"data_source": "image",
|
|
||||||
"args": {
|
|
||||||
"backbone": "refugelike",
|
|
||||||
"freeze_ratio": 0.0,
|
|
||||||
"augment": true
|
|
||||||
},
|
|
||||||
"warmup_epochs": 0
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "cd",
|
|
||||||
"module": "v4.classes.towers.clinical_tower",
|
|
||||||
"class": "ClinicalEncoder",
|
|
||||||
"data_source": "matrix",
|
|
||||||
"args": {
|
|
||||||
"hidden_dim": 128
|
|
||||||
},
|
|
||||||
"warmup_epochs": 40,
|
|
||||||
"warmup_exclude_towers": [
|
|
||||||
"img"
|
|
||||||
]
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"bridge": {
|
|
||||||
"mode": "embedding_mlp",
|
|
||||||
"fusion_dim": 256,
|
|
||||||
"hidden_dim": 256
|
|
||||||
},
|
|
||||||
"training": {
|
|
||||||
"lr": 0.0001,
|
|
||||||
"batch_size": 16,
|
|
||||||
"bcd_prob": 0.5,
|
|
||||||
"warmup_tower_epochs": 3,
|
|
||||||
"warmup_fused_epochs": 3,
|
|
||||||
"tune_binary_threshold": true
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"mean_val_auc": 0.9036764705882353,
|
|
||||||
"std_val_auc": 0.02318218054475656,
|
|
||||||
"mean_test_auc": 0.9007352941176471,
|
|
||||||
"std_test_auc": 0.04178903599524356,
|
|
||||||
"elapsed_s": 1563.1,
|
|
||||||
"fold_results": [
|
|
||||||
{
|
|
||||||
"fold": 0,
|
|
||||||
"val_auc": 0.9044117647058824,
|
|
||||||
"val_acc": 0.8571428571428571,
|
|
||||||
"val_n": 42,
|
|
||||||
"val_kappa": 0.42727272727272725,
|
|
||||||
"val_mcc": 0.4622975667767223,
|
|
||||||
"val_f1": 0.7083333333333333,
|
|
||||||
"val_threshold": 0.1659889668226242,
|
|
||||||
"test_auc": 0.8566176470588236,
|
|
||||||
"test_acc": 0.9285714285714286,
|
|
||||||
"test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 1,
|
|
||||||
"val_auc": 0.8970588235294118,
|
|
||||||
"val_acc": 0.8809523809523809,
|
|
||||||
"val_n": 42,
|
|
||||||
"val_kappa": 0.5945945945945946,
|
|
||||||
"val_mcc": 0.5965587590013045,
|
|
||||||
"val_f1": 0.7971014492753623,
|
|
||||||
"val_threshold": 0.017083797603845596,
|
|
||||||
"test_auc": 0.9044117647058824,
|
|
||||||
"test_acc": 0.8095238095238095,
|
|
||||||
"test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 2,
|
|
||||||
"val_auc": 0.863970588235294,
|
|
||||||
"val_acc": 0.8333333333333334,
|
|
||||||
"val_n": 42,
|
|
||||||
"val_kappa": 0.43243243243243246,
|
|
||||||
"val_mcc": 0.4338609156373123,
|
|
||||||
"val_f1": 0.7159420289855072,
|
|
||||||
"val_threshold": 0.09843172132968903,
|
|
||||||
"test_auc": 0.9044117647058824,
|
|
||||||
"test_acc": 0.8571428571428571,
|
|
||||||
"test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 3,
|
|
||||||
"val_auc": 0.9227941176470588,
|
|
||||||
"val_acc": 0.8809523809523809,
|
|
||||||
"val_n": 42,
|
|
||||||
"val_kappa": 0.631578947368421,
|
|
||||||
"val_mcc": 0.6333004963811236,
|
|
||||||
"val_f1": 0.8156277436347674,
|
|
||||||
"val_threshold": 0.01531070377677679,
|
|
||||||
"test_auc": 0.8639705882352942,
|
|
||||||
"test_acc": 0.8571428571428571,
|
|
||||||
"test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 4,
|
|
||||||
"val_auc": 0.9301470588235294,
|
|
||||||
"val_acc": 0.9047619047619048,
|
|
||||||
"val_n": 42,
|
|
||||||
"val_kappa": 0.6181818181818182,
|
|
||||||
"val_mcc": 0.6688560540599386,
|
|
||||||
"val_f1": 0.8055555555555556,
|
|
||||||
"val_threshold": 0.19587548077106476,
|
|
||||||
"test_auc": 0.9742647058823529,
|
|
||||||
"test_acc": 0.9285714285714286,
|
|
||||||
"test_n": 42
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
@@ -1,264 +0,0 @@
|
|||||||
{
|
|
||||||
"run_name": "v4/ensemble_fused",
|
|
||||||
"eval_stage": "hb",
|
|
||||||
"config": {
|
|
||||||
"_notes": [
|
|
||||||
"V4 stage-pipeline: warm \u2192 fusion stages with parallel head stages.",
|
|
||||||
"Bridges are pure embedding producers; heads are separate swappable stages.",
|
|
||||||
"BCD-eligible heads are sampled during tower_warmup and main phases.",
|
|
||||||
"eval_stage names the fusion stage whose primary head is used for final metrics."
|
|
||||||
],
|
|
||||||
"run_name": "v4/ensemble_fused",
|
|
||||||
"num_classes": 2,
|
|
||||||
"label_filter": [
|
|
||||||
0,
|
|
||||||
1
|
|
||||||
],
|
|
||||||
"split_identity_level": 1,
|
|
||||||
"eval_stage": "hb",
|
|
||||||
"save_predictions": false,
|
|
||||||
"seed": 1234,
|
|
||||||
"folds": 5,
|
|
||||||
"fold_seed": 100,
|
|
||||||
"output_root": "v4/results",
|
|
||||||
"out_dir_tags": [
|
|
||||||
"binary",
|
|
||||||
"ntower"
|
|
||||||
],
|
|
||||||
"data": {
|
|
||||||
"module": "v4.classes.profiles.v4papila",
|
|
||||||
"args": {
|
|
||||||
"image_dir": "Papila/FundusImages",
|
|
||||||
"clinical_dir": "Papila/ClinicalData",
|
|
||||||
"label_col": "Diagnosis",
|
|
||||||
"iop_corr_method": "ratio",
|
|
||||||
"iop_drop_raw": true,
|
|
||||||
"exclude_cols": [
|
|
||||||
"Axial_Length"
|
|
||||||
],
|
|
||||||
"in_memory_cache": true
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"towers": [
|
|
||||||
{
|
|
||||||
"name": "img",
|
|
||||||
"module": "v4.classes.towers.image_tower",
|
|
||||||
"class": "ImageEncoder",
|
|
||||||
"data_source": "image",
|
|
||||||
"args": {
|
|
||||||
"backbone": "refugelike",
|
|
||||||
"freeze_ratio": 0.0,
|
|
||||||
"augment": true
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "cd",
|
|
||||||
"module": "v4.classes.towers.clinical_tower",
|
|
||||||
"class": "ClinicalEncoder",
|
|
||||||
"data_source": "matrix",
|
|
||||||
"args": {
|
|
||||||
"hidden_dim": 128
|
|
||||||
}
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"stages": [
|
|
||||||
{
|
|
||||||
"name": "cd_warm",
|
|
||||||
"type": "warm",
|
|
||||||
"tower": "cd",
|
|
||||||
"level": "eye",
|
|
||||||
"epochs": 40
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "img_aux",
|
|
||||||
"type": "head",
|
|
||||||
"input": "img",
|
|
||||||
"train_with": "nt",
|
|
||||||
"bcd": true
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "cd_aux",
|
|
||||||
"type": "head",
|
|
||||||
"input": "cd",
|
|
||||||
"train_with": "nt",
|
|
||||||
"bcd": true
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "nt",
|
|
||||||
"type": "fusion",
|
|
||||||
"module": "v4.classes.bridges.fusion_bridge",
|
|
||||||
"class": "FusionBridge",
|
|
||||||
"inputs": [
|
|
||||||
"img",
|
|
||||||
"cd"
|
|
||||||
],
|
|
||||||
"level": "eye",
|
|
||||||
"epochs": 36,
|
|
||||||
"train_towers": true,
|
|
||||||
"warmup": {
|
|
||||||
"tower_epochs": 3,
|
|
||||||
"fused_epochs": 3
|
|
||||||
},
|
|
||||||
"args": {
|
|
||||||
"fusion_dim": 256
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "nt_head",
|
|
||||||
"type": "head",
|
|
||||||
"input": "nt",
|
|
||||||
"train_with": "nt"
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "hb",
|
|
||||||
"type": "fusion",
|
|
||||||
"module": "v4.classes.bridges.hyperbridge",
|
|
||||||
"class": "HyperBridge",
|
|
||||||
"inputs": {
|
|
||||||
"a": "nt",
|
|
||||||
"b": "nt"
|
|
||||||
},
|
|
||||||
"level": "patient",
|
|
||||||
"epochs": 10,
|
|
||||||
"args": {
|
|
||||||
"hidden_dim": 256,
|
|
||||||
"mode": "embedding_mlp"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "hb_head",
|
|
||||||
"type": "head",
|
|
||||||
"input": "hb",
|
|
||||||
"train_with": "hb"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"training": {
|
|
||||||
"lr": 0.0001,
|
|
||||||
"batch_size": 16,
|
|
||||||
"bcd_prob": 0.5,
|
|
||||||
"tune_binary_threshold": true
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"mean_val_auc": 0.8926470588235293,
|
|
||||||
"std_val_auc": 0.06105155516072669,
|
|
||||||
"mean_test_auc": 0.9022058823529413,
|
|
||||||
"std_test_auc": 0.035001853633564395,
|
|
||||||
"elapsed_s": 1525.9,
|
|
||||||
"fold_results": [
|
|
||||||
{
|
|
||||||
"fold": 0,
|
|
||||||
"nt_val_auc": 0.858647936786655,
|
|
||||||
"nt_val_acc": 0.8928571428571429,
|
|
||||||
"nt_val_n": 84,
|
|
||||||
"nt_val_kappa": 0.6272189349112426,
|
|
||||||
"nt_val_mcc": 0.6411186083279721,
|
|
||||||
"nt_val_f1": 0.8124534854874721,
|
|
||||||
"nt_val_threshold": 0.012987074442207813,
|
|
||||||
"nt_test_auc": 0.9157155399473222,
|
|
||||||
"nt_test_acc": 0.9047619047619048,
|
|
||||||
"nt_test_n": 84,
|
|
||||||
"hb_val_auc": 0.9080882352941178,
|
|
||||||
"hb_val_acc": 0.8333333333333334,
|
|
||||||
"hb_val_n": 42,
|
|
||||||
"hb_val_kappa": 0.2898550724637682,
|
|
||||||
"hb_val_mcc": 0.33633639699815626,
|
|
||||||
"hb_val_f1": 0.6338729763387297,
|
|
||||||
"hb_val_threshold": 0.011891158297657967,
|
|
||||||
"hb_test_auc": 0.9522058823529411,
|
|
||||||
"hb_test_acc": 0.9523809523809523,
|
|
||||||
"hb_test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 1,
|
|
||||||
"nt_val_auc": 0.8282828282828283,
|
|
||||||
"nt_val_acc": 0.8928571428571429,
|
|
||||||
"nt_val_n": 84,
|
|
||||||
"nt_val_kappa": 0.6111111111111112,
|
|
||||||
"nt_val_mcc": 0.6633249580710799,
|
|
||||||
"nt_val_f1": 0.801418439716312,
|
|
||||||
"nt_val_threshold": 0.008082838729023933,
|
|
||||||
"nt_test_auc": 0.8726953467954346,
|
|
||||||
"nt_test_acc": 0.8690476190476191,
|
|
||||||
"nt_test_n": 84,
|
|
||||||
"hb_val_auc": 0.7904411764705882,
|
|
||||||
"hb_val_acc": 0.9047619047619048,
|
|
||||||
"hb_val_n": 42,
|
|
||||||
"hb_val_kappa": 0.6181818181818182,
|
|
||||||
"hb_val_mcc": 0.6688560540599386,
|
|
||||||
"hb_val_f1": 0.8055555555555556,
|
|
||||||
"hb_val_threshold": 0.013159404508769512,
|
|
||||||
"hb_test_auc": 0.8970588235294117,
|
|
||||||
"hb_test_acc": 0.8571428571428571,
|
|
||||||
"hb_test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 2,
|
|
||||||
"nt_val_auc": 0.8906882591093117,
|
|
||||||
"nt_val_acc": 0.8452380952380952,
|
|
||||||
"nt_val_n": 84,
|
|
||||||
"nt_val_kappa": 0.5125,
|
|
||||||
"nt_val_mcc": 0.5217535056401378,
|
|
||||||
"nt_val_f1": 0.7548821548821549,
|
|
||||||
"nt_val_threshold": 0.052708517760038376,
|
|
||||||
"nt_test_auc": 0.856060606060606,
|
|
||||||
"nt_test_acc": 0.8333333333333334,
|
|
||||||
"nt_test_n": 84,
|
|
||||||
"hb_val_auc": 0.9117647058823529,
|
|
||||||
"hb_val_acc": 0.8333333333333334,
|
|
||||||
"hb_val_n": 42,
|
|
||||||
"hb_val_kappa": 0.36909871244635195,
|
|
||||||
"hb_val_mcc": 0.3833788364965519,
|
|
||||||
"hb_val_f1": 0.6814734561213435,
|
|
||||||
"hb_val_threshold": 0.0401872955262661,
|
|
||||||
"hb_test_auc": 0.8566176470588236,
|
|
||||||
"hb_test_acc": 0.8571428571428571,
|
|
||||||
"hb_test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 3,
|
|
||||||
"nt_val_auc": 0.921875,
|
|
||||||
"nt_val_acc": 0.8928571428571429,
|
|
||||||
"nt_val_n": 84,
|
|
||||||
"nt_val_kappa": 0.6758147512864494,
|
|
||||||
"nt_val_mcc": 0.6797955088067001,
|
|
||||||
"nt_val_f1": 0.837593984962406,
|
|
||||||
"nt_val_threshold": 0.708580732345581,
|
|
||||||
"nt_test_auc": 0.8631578947368421,
|
|
||||||
"nt_test_acc": 0.8452380952380952,
|
|
||||||
"nt_test_n": 84,
|
|
||||||
"hb_val_auc": 0.9779411764705882,
|
|
||||||
"hb_val_acc": 0.9285714285714286,
|
|
||||||
"hb_val_n": 42,
|
|
||||||
"hb_val_kappa": 0.7567567567567568,
|
|
||||||
"hb_val_mcc": 0.7592566023652966,
|
|
||||||
"hb_val_f1": 0.8782608695652174,
|
|
||||||
"hb_val_threshold": 0.04653067886829376,
|
|
||||||
"hb_test_auc": 0.875,
|
|
||||||
"hb_test_acc": 0.8809523809523809,
|
|
||||||
"hb_test_n": 42
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"fold": 4,
|
|
||||||
"nt_val_auc": 0.8279192273924496,
|
|
||||||
"nt_val_acc": 0.8571428571428571,
|
|
||||||
"nt_val_n": 84,
|
|
||||||
"nt_val_kappa": 0.46325878594249204,
|
|
||||||
"nt_val_mcc": 0.49610717544581684,
|
|
||||||
"nt_val_f1": 0.7269772481040087,
|
|
||||||
"nt_val_threshold": 0.06419441103935242,
|
|
||||||
"nt_test_auc": 0.8961397058823529,
|
|
||||||
"nt_test_acc": 0.8928571428571429,
|
|
||||||
"nt_test_n": 84,
|
|
||||||
"hb_val_auc": 0.875,
|
|
||||||
"hb_val_acc": 0.9047619047619048,
|
|
||||||
"hb_val_n": 42,
|
|
||||||
"hb_val_kappa": 0.6181818181818182,
|
|
||||||
"hb_val_mcc": 0.6688560540599386,
|
|
||||||
"hb_val_f1": 0.8055555555555556,
|
|
||||||
"hb_val_threshold": 0.23509082198143005,
|
|
||||||
"hb_test_auc": 0.9301470588235294,
|
|
||||||
"hb_test_acc": 0.8809523809523809,
|
|
||||||
"hb_test_n": 42
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
[
|
||||||
|
{
|
||||||
|
"run_name": "experiments/ensemble_fused/geom_gt"
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
[
|
||||||
|
{
|
||||||
|
"_note": "Geometry vector injection, GT source, all 5 features — v4 equivalent of v3 phase6a.",
|
||||||
|
"run_name": "experiments/geometry_vec_gt/dim5"
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
[
|
||||||
|
{
|
||||||
|
"run_name": "experiments/ensemble_fused/no_geom"
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
[
|
||||||
|
{
|
||||||
|
"_note": "Replacement run of ensemble_fused (img + cd) — keeps original alongside.",
|
||||||
|
"run_name": "experiments/tri_v1/baseline_ensemble"
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
[
|
||||||
|
{
|
||||||
|
"_note": "Standalone geometry tower (UNet seg, MonoBridge) — 10-rep baseline.",
|
||||||
|
"run_name": "experiments/tri_v1/baseline_solo"
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
[
|
||||||
|
{
|
||||||
|
"_note": "Tritower with default params (bcd_prob=0.5, cw=false, nt_epochs=36) — 10-rep baseline.",
|
||||||
|
"run_name": "experiments/tri_v1/baseline_tri"
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
[
|
||||||
|
{ "_note": "tritower grid pilot (3 reps each) — 3 bcd × 2 cw × 3 nt_epochs = 18 cells", "run_name": "experiments/tri_v1/grid/bcd35_cw0_nt15", "reps": 3, "overrides": { "training": { "bcd_prob": 0.35, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 15 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd35_cw0_nt25", "reps": 3, "overrides": { "training": { "bcd_prob": 0.35, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 25 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd35_cw0_nt36", "reps": 3, "overrides": { "training": { "bcd_prob": 0.35, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 36 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd35_cw1_nt15", "reps": 3, "overrides": { "training": { "bcd_prob": 0.35, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 15 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd35_cw1_nt25", "reps": 3, "overrides": { "training": { "bcd_prob": 0.35, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 25 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd35_cw1_nt36", "reps": 3, "overrides": { "training": { "bcd_prob": 0.35, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 36 } } },
|
||||||
|
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd50_cw0_nt15", "reps": 3, "overrides": { "training": { "bcd_prob": 0.50, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 15 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd50_cw0_nt25", "reps": 3, "overrides": { "training": { "bcd_prob": 0.50, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 25 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd50_cw0_nt36", "reps": 3, "overrides": { "training": { "bcd_prob": 0.50, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 36 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd50_cw1_nt15", "reps": 3, "overrides": { "training": { "bcd_prob": 0.50, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 15 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd50_cw1_nt25", "reps": 3, "overrides": { "training": { "bcd_prob": 0.50, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 25 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd50_cw1_nt36", "reps": 3, "overrides": { "training": { "bcd_prob": 0.50, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 36 } } },
|
||||||
|
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd75_cw0_nt15", "reps": 3, "overrides": { "training": { "bcd_prob": 0.75, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 15 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd75_cw0_nt25", "reps": 3, "overrides": { "training": { "bcd_prob": 0.75, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 25 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd75_cw0_nt36", "reps": 3, "overrides": { "training": { "bcd_prob": 0.75, "class_weighted": false } }, "stage_overrides": { "nt": { "epochs": 36 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd75_cw1_nt15", "reps": 3, "overrides": { "training": { "bcd_prob": 0.75, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 15 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd75_cw1_nt25", "reps": 3, "overrides": { "training": { "bcd_prob": 0.75, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 25 } } },
|
||||||
|
{ "run_name": "experiments/tri_v1/grid/bcd75_cw1_nt36", "reps": 3, "overrides": { "training": { "bcd_prob": 0.75, "class_weighted": true } }, "stage_overrides": { "nt": { "epochs": 36 } } }
|
||||||
|
]
|
||||||
Reference in New Issue
Block a user