Refactor geometry feature loaders and enhance distributed client status reporting

- Updated GTGeometryLoader to streamline geometry vector computation and caching.
- Introduced UNetGeometryLoader for UNet-derived geometry vectors.
- Added sample collection method in PapilaBundle for better data handling.
- Refined GeometrySegEncoder and ImageEncoder to utilize new sample collection.
- Enhanced distributed client with heartbeat mechanism for improved job tracking.
- Added new configuration files for UNet-derived geometry integration.
This commit is contained in:
rpotter6298
2026-04-29 11:05:19 +02:00
parent 512ebd13b2
commit af813bbb62
11 changed files with 467 additions and 153 deletions
+17 -3
View File
@@ -151,11 +151,12 @@ class ImageEncoder(TowerBase):
def early_pass(self, context) -> None:
"""Per-fold setup: warm tensor cache (if enabled), publish geometry vectors."""
data = context.require("data")
data = context.require("data")
split = context.require("split")
if self._cache_transformed:
self._tensor_cache.clear()
n = self._warm_tensor_cache(data, context.require("split"))
n = self._warm_tensor_cache(data, split)
print(
f"[ImageEncoder] warmed transformed-tensor cache for {n} entries "
f"({self._name})",
@@ -164,7 +165,20 @@ class ImageEncoder(TowerBase):
if self._geom_loader is None:
return
self._geom_loader.precompute(data.df, patient_col=data.patient_col)
train_samples = data.collect_samples(split.train)
all_samples = train_samples + data.collect_samples(split.val)
if split.test is not None:
all_samples += data.collect_samples(split.test)
if hasattr(self._geom_loader, "reset_cache"):
self._geom_loader.reset_cache()
if hasattr(self._geom_loader, "reset_weights"):
self._geom_loader.reset_weights()
if hasattr(self._geom_loader, "finetune"):
self._geom_loader.finetune(train_samples)
self._geom_loader.precompute(all_samples)
vecs = self._geom_loader.all_vectors()
context.put(self.EPC_GEOMETRY_KEY, vecs)
print(