Add fail-closed source class and SAM fallback filters
This commit is contained in:
@@ -108,6 +108,7 @@ def main() -> int:
|
||||
parser.add_argument("--min-dimension-ratio", type=float, default=0.5)
|
||||
parser.add_argument("--max-dimension-ratio", type=float, default=2.0)
|
||||
parser.add_argument("--max-prompts-per-pass", type=int, default=96)
|
||||
parser.add_argument("--fallback-policy", choices=("retain", "drop"), default="retain")
|
||||
parser.add_argument("--force", action="store_true")
|
||||
args = parser.parse_args()
|
||||
if args.output_dir.exists():
|
||||
@@ -134,7 +135,7 @@ def main() -> int:
|
||||
with Image.open(source_image) as image:
|
||||
width, height = image.size
|
||||
source_boxes = read_boxes(source_label, width, height)
|
||||
output_boxes = list(source_boxes)
|
||||
output_boxes: list[tuple[float, float, float, float] | None] = list(source_boxes)
|
||||
if source_boxes:
|
||||
for start in range(0, len(source_boxes), args.max_prompts_per_pass):
|
||||
source_chunk = source_boxes[start : start + args.max_prompts_per_pass]
|
||||
@@ -153,6 +154,8 @@ def main() -> int:
|
||||
if candidate_index < 0 or overlap < args.min_source_iou:
|
||||
reason_counts["unmatched_mask"] = reason_counts.get("unmatched_mask", 0) + 1
|
||||
fallback_count += 1
|
||||
if args.fallback_policy == "drop":
|
||||
output_boxes[start + local_index] = None
|
||||
continue
|
||||
candidate = candidates[candidate_index]
|
||||
if plausible_refinement(
|
||||
@@ -171,14 +174,22 @@ def main() -> int:
|
||||
else:
|
||||
reason_counts["geometry_gate"] = reason_counts.get("geometry_gate", 0) + 1
|
||||
fallback_count += 1
|
||||
if args.fallback_policy == "drop":
|
||||
output_boxes[start + local_index] = None
|
||||
del result, masks
|
||||
if model.predictor is not None:
|
||||
model.predictor.reset_image()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
target_label.write_text("\n".join(yolo_line(box, width, height) for box in output_boxes) + ("\n" if output_boxes else ""), encoding="utf-8")
|
||||
retained_boxes = [box for box in output_boxes if box is not None]
|
||||
target_label.write_text("\n".join(yolo_line(box, width, height) for box in retained_boxes) + ("\n" if retained_boxes else ""), encoding="utf-8")
|
||||
output_tile = dict(tile)
|
||||
output_tile.update({"image_path": str(target_image), "label_path": str(target_label)})
|
||||
output_tile.update({
|
||||
"image_path": str(target_image),
|
||||
"label_path": str(target_label),
|
||||
"label_count": len(retained_boxes),
|
||||
"is_negative": not retained_boxes,
|
||||
})
|
||||
output_tiles.append(output_tile)
|
||||
print(f"{index}/{len(summary['tiles'])} {tile['sample_slug']}: {len(source_boxes)}", flush=True)
|
||||
|
||||
@@ -188,7 +199,14 @@ def main() -> int:
|
||||
"output_dir": str(args.output_dir),
|
||||
"dataset_yaml": str(args.output_dir / "dataset.yaml"),
|
||||
"tiles": output_tiles,
|
||||
"label_semantics": "sam_visible_roof_with_official_footprint_fallback",
|
||||
"label_semantics": (
|
||||
"sam_visible_roof_only"
|
||||
if args.fallback_policy == "drop"
|
||||
else "sam_visible_roof_with_official_footprint_fallback"
|
||||
),
|
||||
"label_count": refined_count if args.fallback_policy == "drop" else summary["label_count"],
|
||||
"positive_tile_count": sum(not tile["is_negative"] for tile in output_tiles),
|
||||
"negative_tile_count": sum(tile["is_negative"] for tile in output_tiles),
|
||||
}
|
||||
)
|
||||
summary_path = args.output_dir / "yolo_tile_dataset_summary.json"
|
||||
@@ -213,8 +231,10 @@ def main() -> int:
|
||||
"min_dimension_ratio": args.min_dimension_ratio,
|
||||
"max_dimension_ratio": args.max_dimension_ratio,
|
||||
"max_prompts_per_pass": args.max_prompts_per_pass,
|
||||
"fallback_policy": args.fallback_policy,
|
||||
"refined_label_count": refined_count,
|
||||
"fallback_label_count": fallback_count,
|
||||
"dropped_fallback_label_count": fallback_count if args.fallback_policy == "drop" else 0,
|
||||
"fallback_reason_counts": reason_counts,
|
||||
}
|
||||
(args.output_dir / "sam-refinement.json").write_text(json.dumps(evidence, indent=2), encoding="utf-8")
|
||||
|
||||
Reference in New Issue
Block a user