pre-restructure

This commit is contained in:
rpotter6298
2026-02-26 06:47:58 +01:00
parent 17a255daa8
commit 8980cf5f9b
10 changed files with 977 additions and 63 deletions
+30
View File
@@ -401,6 +401,7 @@ class RefugeClassification:
f"[classifier] Building datasets from {len(candidates)} labelled samples (train/val)"
)
skipped: List[str] = []
for sample in tqdm(
candidates,
desc="Preparing records",
@@ -410,6 +411,7 @@ class RefugeClassification:
try:
geom, disc_mask, cup_mask = self._resolve_geometry(sample, crop_scale)
except RuntimeError:
skipped.append(sample.sample_id)
continue
record = RefugeClassificationRecord(
sample=sample,
@@ -424,6 +426,12 @@ class RefugeClassification:
else:
val_records.append(record)
if skipped:
print(
f"[classifier] WARNING: {len(skipped)}/{len(candidates)} samples skipped "
f"due to empty segmentation mask: {skipped}"
)
if (not val_records or self.use_all_labeled) and train_records and self.auto_val_ratio > 0.0:
rng = random.Random(42)
label_groups: Dict[int, List[RefugeClassificationRecord]] = {}
@@ -578,17 +586,20 @@ class RefugeClassification:
) -> List[RefugeClassificationRecord]:
scale = crop_scale if crop_scale is not None else self.crop_scale
records: List[RefugeClassificationRecord] = []
skipped: List[str] = []
iterator: Iterable[RefugeSample]
if progress_prefix is not None:
iterator = tqdm(samples, desc=progress_prefix, unit="sample", leave=False)
else:
iterator = samples
labeled = [s for s in samples if s.label is not None]
for sample in iterator:
if sample.label is None:
continue
try:
geom, disc_mask, cup_mask = self._resolve_geometry(sample, scale)
except RuntimeError:
skipped.append(sample.sample_id)
continue
records.append(
RefugeClassificationRecord(
@@ -598,8 +609,27 @@ class RefugeClassification:
cup_mask=cup_mask,
)
)
prefix = f"[{progress_prefix}]" if progress_prefix else "[classifier]"
if skipped:
print(
f"{prefix} WARNING: {len(skipped)}/{len(labeled)} samples skipped "
f"due to empty segmentation mask: {skipped}"
)
else:
print(f"{prefix} All {len(labeled)} samples processed successfully.")
return records
def clear_disk_cache(self) -> None:
"""Delete all cached geometry/mask .npz files in cache_dir."""
if self.cache_dir is None or not self.cache_dir.exists():
return
removed = 0
for f in self.cache_dir.glob("*.npz"):
f.unlink()
removed += 1
self.geometry_cache.clear()
print(f"[classifier] Cleared {removed} cached geometry files from {self.cache_dir}")
def _cache_key(self, sample_id: str, scale: float) -> str:
scale_tag = int(round(scale * 100))
return f"{sample_id}_s{scale_tag}"