mirror of
https://github.com/firestar5683/StarPilot.git
synced 2026-09-30 03:13:48 +08:00
Analytical Techniques
This commit is contained in:
@@ -20,6 +20,19 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--output", type=Path, required=True, help="New isolated dataset root.")
|
||||
parser.add_argument("--positive-manifest", type=Path, action="append", default=[], help="Reviewed positive crop manifest. Repeat as needed.")
|
||||
parser.add_argument("--reject-manifest", type=Path, action="append", default=[], help="Reviewed classifier reject manifest. Repeat as needed.")
|
||||
parser.add_argument(
|
||||
"--exclude-base-record-key",
|
||||
action="append",
|
||||
default=[],
|
||||
help="Remove inherited samples whose staged filename contains this corrected record key. Repeat as needed.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--repeat-reject-record",
|
||||
action="append",
|
||||
default=[],
|
||||
metavar="RECORD_KEY=COUNT",
|
||||
help="Stage a reviewed reject COUNT times to give a corrected hard negative more training weight.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--advisory-as-reject",
|
||||
action="store_true",
|
||||
@@ -76,12 +89,45 @@ def keep_advisory_reject(row: dict[str, str], fraction: float) -> bool:
|
||||
return int.from_bytes(digest[:8], "big") / 2**64 < fraction
|
||||
|
||||
|
||||
def safe_record_key(record_key: str) -> str:
|
||||
return "".join(char if char.isalnum() or char in "._-" else "_" for char in record_key)[:100]
|
||||
|
||||
|
||||
def remove_inherited_records(root: Path, record_keys: list[str]) -> int:
|
||||
safe_keys = tuple(filter(None, (safe_record_key(record_key) for record_key in record_keys)))
|
||||
if not safe_keys:
|
||||
return 0
|
||||
removed = 0
|
||||
for split in ("train", "val"):
|
||||
for path in (root / split).rglob("*"):
|
||||
if path.is_file() and any(record_key in path.name for record_key in safe_keys):
|
||||
path.unlink()
|
||||
removed += 1
|
||||
return removed
|
||||
|
||||
|
||||
def parse_reject_repeat_counts(specs: list[str]) -> dict[str, int]:
|
||||
repeat_counts: dict[str, int] = {}
|
||||
for spec in specs:
|
||||
record_key, separator, count_text = spec.rpartition("=")
|
||||
if not separator or not record_key:
|
||||
raise ValueError(f"Invalid --repeat-reject-record value: {spec!r}")
|
||||
try:
|
||||
count = int(count_text)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"Invalid reject repeat count: {spec!r}") from exc
|
||||
if count < 1:
|
||||
raise ValueError(f"Reject repeat count must be at least 1: {spec!r}")
|
||||
repeat_counts[record_key] = count
|
||||
return repeat_counts
|
||||
|
||||
|
||||
def stage_crop(source: Path, destination_dir: Path, record_key: str) -> bool:
|
||||
if not source.is_file():
|
||||
return False
|
||||
digest = hashlib.sha256(source.read_bytes()).hexdigest()[:16]
|
||||
suffix = source.suffix.lower() if source.suffix.lower() in (".jpg", ".jpeg", ".png") else ".jpg"
|
||||
safe_key = "".join(char if char.isalnum() or char in "._-" else "_" for char in record_key)[:100]
|
||||
safe_key = safe_record_key(record_key)
|
||||
destination_dir.mkdir(parents=True, exist_ok=True)
|
||||
destination = destination_dir / f"review_{safe_key}_{digest}{suffix}"
|
||||
if not destination.exists():
|
||||
@@ -103,6 +149,8 @@ def main() -> int:
|
||||
raise FileExistsError(f"Output dataset already exists: {output}")
|
||||
shutil.copytree(base, output, copy_function=shutil.copyfile)
|
||||
appledouble_removed = remove_appledouble_files(output)
|
||||
inherited_records_removed = remove_inherited_records(output, args.exclude_base_record_key)
|
||||
reject_repeat_counts = parse_reject_repeat_counts(args.repeat_reject_record)
|
||||
|
||||
positive_counts: Counter[str] = Counter()
|
||||
reject_counts: Counter[str] = Counter()
|
||||
@@ -132,10 +180,16 @@ def main() -> int:
|
||||
for row in read_rows(args.reject_manifest):
|
||||
split = row.get("split", "")
|
||||
source = Path(row.get("crop_path", "")).expanduser()
|
||||
if split not in ("train", "val") or not stage_crop(source, output / split / "reject", row.get("record_key", "reject")):
|
||||
record_key = row.get("record_key", "reject")
|
||||
repeat_count = reject_repeat_counts.get(record_key, 1) if split == "train" else 1
|
||||
staged = split in ("train", "val")
|
||||
for repeat_index in range(repeat_count):
|
||||
staged_key = record_key if repeat_index == 0 else f"{record_key}_repeat_{repeat_index:03d}"
|
||||
staged = staged and stage_crop(source, output / split / "reject", staged_key)
|
||||
if not staged:
|
||||
skipped += 1
|
||||
continue
|
||||
reject_counts[split] += 1
|
||||
reject_counts[split] += repeat_count
|
||||
|
||||
appledouble_removed += remove_appledouble_files(output)
|
||||
for split in ("train", "val"):
|
||||
@@ -150,6 +204,7 @@ def main() -> int:
|
||||
"reject_counts": dict(sorted(reject_counts.items())),
|
||||
"skipped": skipped,
|
||||
"appledouble_removed": appledouble_removed,
|
||||
"inherited_records_removed": inherited_records_removed,
|
||||
}
|
||||
summary_path = output / "review_dataset_summary.json"
|
||||
summary_path.write_text(json.dumps(summary, indent=2, sort_keys=True) + "\n", encoding="utf-8")
|
||||
|
||||
@@ -305,6 +305,58 @@ def classifier_reject_row(row: dict[str, str], split: str) -> dict[str, object]:
|
||||
}
|
||||
|
||||
|
||||
RUNTIME_REJECT_CROP_EXPANSIONS = (
|
||||
(0.00, 0.00, 0.00, 0.00),
|
||||
(0.10, 0.06, 0.10, 0.12),
|
||||
(0.00, 0.00, 0.18, 0.18),
|
||||
)
|
||||
|
||||
|
||||
def classifier_reject_variant_rows(
|
||||
row: dict[str, str],
|
||||
split: str,
|
||||
output_dir: Path,
|
||||
overwrite: bool,
|
||||
) -> list[dict[str, object]]:
|
||||
rows = [classifier_reject_row(row, split)]
|
||||
if row.get("review_ignore_reason") != "conditional_restriction":
|
||||
return rows
|
||||
|
||||
frame_path = Path(row.get("frame_path", "")).expanduser()
|
||||
frame = cv2.imread(str(frame_path))
|
||||
bbox = parse_bbox(row.get("review_bbox") or row.get("bbox", ""))
|
||||
if frame is None or bbox is None:
|
||||
raise RuntimeError(f"Cannot generate conditional reject crops for {row['record_key']}: unreadable frame or bbox")
|
||||
|
||||
image_h, image_w = frame.shape[:2]
|
||||
x1, y1, x2, y2 = bbox
|
||||
box_width = x2 - x1
|
||||
box_height = y2 - y1
|
||||
reject_dir = output_dir / "corrected_classifier_reject_crops"
|
||||
ensure_dir(reject_dir)
|
||||
for index, (expand_left, expand_top, expand_right, expand_bottom) in enumerate(RUNTIME_REJECT_CROP_EXPANSIONS):
|
||||
crop_bbox = (
|
||||
max(int(x1 - box_width * expand_left), 0),
|
||||
max(int(y1 - box_height * expand_top), 0),
|
||||
min(int(x2 + box_width * expand_right), image_w),
|
||||
min(int(y2 + box_height * expand_bottom), image_h),
|
||||
)
|
||||
crop_x1, crop_y1, crop_x2, crop_y2 = crop_bbox
|
||||
crop = frame[crop_y1:crop_y2, crop_x1:crop_x2]
|
||||
crop_path = reject_dir / f"{safe_stem(row['record_key'])}_runtime_expansion_{index}.jpg"
|
||||
if crop.size == 0:
|
||||
raise RuntimeError(f"Cannot generate conditional reject crop for {row['record_key']}: empty bbox {crop_bbox}")
|
||||
if overwrite or not crop_path.is_file():
|
||||
if not cv2.imwrite(str(crop_path), crop, [cv2.IMWRITE_JPEG_QUALITY, 94]):
|
||||
raise RuntimeError(f"Cannot write conditional reject crop for {row['record_key']}: {crop_path}")
|
||||
variant = classifier_reject_row(row, split)
|
||||
variant["record_key"] = f"{row['record_key']}_runtime_expansion_{index}"
|
||||
variant["crop_path"] = str(crop_path)
|
||||
variant["crop_bbox"] = ",".join(str(value) for value in crop_bbox)
|
||||
rows.append(variant)
|
||||
return rows
|
||||
|
||||
|
||||
def corrected_classifier_crop(
|
||||
row: dict[str, str],
|
||||
output_dir: Path,
|
||||
@@ -508,7 +560,7 @@ def main() -> int:
|
||||
|
||||
for row in classifier_reject_rows:
|
||||
split = split_for_key(split_group_key(row), args.val_modulo, args.val_remainder)
|
||||
reject_rows.append(classifier_reject_row(row, split))
|
||||
reject_rows.extend(classifier_reject_variant_rows(row, split, output_dir, args.overwrite))
|
||||
|
||||
write_csv(classifier_manifest, CLASSIFIER_FIELDNAMES, classifier_rows)
|
||||
write_csv(runtime_manifest, RUNTIME_FIELDNAMES, runtime_rows)
|
||||
|
||||
@@ -235,3 +235,53 @@ def test_corrected_bbox_requires_readable_source_frame(tmp_path):
|
||||
|
||||
with pytest.raises(RuntimeError, match="unreadable frame"):
|
||||
import_queue.corrected_classifier_crop(row, tmp_path, overwrite=False)
|
||||
|
||||
|
||||
def test_corrected_record_removes_inherited_classifier_sample(tmp_path):
|
||||
stale = tmp_path / "train" / "55" / "base_review_bad_record_key_hash.jpg"
|
||||
retained = tmp_path / "train" / "55" / "base_review_other_record_hash.jpg"
|
||||
stale.parent.mkdir(parents=True)
|
||||
stale.write_bytes(b"stale")
|
||||
retained.write_bytes(b"retained")
|
||||
|
||||
removed = build_review_classifier.remove_inherited_records(tmp_path, ["bad:record/key"])
|
||||
|
||||
assert removed == 1
|
||||
assert not stale.exists()
|
||||
assert retained.exists()
|
||||
|
||||
|
||||
def test_reject_repeat_spec_preserves_record_key_punctuation():
|
||||
counts = build_review_classifier.parse_reject_repeat_counts(["route/sign=track:55=32"])
|
||||
|
||||
assert counts == {"route/sign=track:55": 32}
|
||||
|
||||
with pytest.raises(ValueError, match="at least 1"):
|
||||
build_review_classifier.parse_reject_repeat_counts(["bad-record=0"])
|
||||
|
||||
|
||||
def test_conditional_reject_generates_runtime_crop_expansions(tmp_path):
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
frame_path = tmp_path / "frame.jpg"
|
||||
crop_path = tmp_path / "crop.jpg"
|
||||
frame = np.zeros((100, 200, 3), dtype=np.uint8)
|
||||
cv2.imwrite(str(frame_path), frame)
|
||||
cv2.imwrite(str(crop_path), frame[20:80, 60:140])
|
||||
row = {
|
||||
"record_key": "conditional-sign",
|
||||
"frame_path": str(frame_path),
|
||||
"crop_path": str(crop_path),
|
||||
"bbox": "60,20,140,80",
|
||||
"review_bbox": "60,20,140,80",
|
||||
"review_ignore_reason": "conditional_restriction",
|
||||
}
|
||||
|
||||
rows = import_queue.classifier_reject_variant_rows(row, "train", tmp_path, overwrite=False)
|
||||
|
||||
assert len(rows) == 4
|
||||
assert rows[1]["crop_bbox"] == "60,20,140,80"
|
||||
assert rows[2]["crop_bbox"] == "52,16,148,87"
|
||||
assert rows[3]["crop_bbox"] == "60,20,154,90"
|
||||
assert all(Path(variant["crop_path"]).is_file() for variant in rows)
|
||||
|
||||
Reference in New Issue
Block a user