pre-restructure
This commit is contained in:
@@ -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}"
|
||||
|
||||
Reference in New Issue
Block a user