Bound proposal-classifier evaluation memory
This commit is contained in:
@@ -100,7 +100,10 @@ def main() -> int:
|
|||||||
parser.add_argument("--proposal-classifier", type=Path)
|
parser.add_argument("--proposal-classifier", type=Path)
|
||||||
parser.add_argument("--proposal-classifier-threshold", type=float, default=0.5)
|
parser.add_argument("--proposal-classifier-threshold", type=float, default=0.5)
|
||||||
parser.add_argument("--proposal-crop-scale", type=float, default=1.4)
|
parser.add_argument("--proposal-crop-scale", type=float, default=1.4)
|
||||||
|
parser.add_argument("--proposal-classifier-batch", type=int, default=64)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
if args.proposal_classifier_batch < 1:
|
||||||
|
parser.error("--proposal-classifier-batch must be positive")
|
||||||
|
|
||||||
from ultralytics import YOLO
|
from ultralytics import YOLO
|
||||||
proposal_classifier = None
|
proposal_classifier = None
|
||||||
@@ -160,8 +163,13 @@ def main() -> int:
|
|||||||
crop = source.crop((left, top, right, bottom))
|
crop = source.crop((left, top, right, bottom))
|
||||||
crops.append(proposal_transform(crop.convert("RGB")))
|
crops.append(proposal_transform(crop.convert("RGB")))
|
||||||
if crops:
|
if crops:
|
||||||
|
probabilities = []
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
probabilities = torch.sigmoid(proposal_classifier(torch.stack(crops).to(args.device)).flatten()).cpu().tolist()
|
for start in range(0, len(crops), args.proposal_classifier_batch):
|
||||||
|
batch = torch.stack(crops[start : start + args.proposal_classifier_batch]).to(args.device)
|
||||||
|
probabilities.extend(
|
||||||
|
torch.sigmoid(proposal_classifier(batch).flatten()).cpu().tolist()
|
||||||
|
)
|
||||||
predictions = [item for item, probability in zip(predictions, probabilities, strict=True) if probability >= args.proposal_classifier_threshold]
|
predictions = [item for item, probability in zip(predictions, probabilities, strict=True) if probability >= args.proposal_classifier_threshold]
|
||||||
observations.append(
|
observations.append(
|
||||||
{
|
{
|
||||||
@@ -208,6 +216,7 @@ def main() -> int:
|
|||||||
"proposal_classifier": str(args.proposal_classifier) if args.proposal_classifier else None,
|
"proposal_classifier": str(args.proposal_classifier) if args.proposal_classifier else None,
|
||||||
"proposal_classifier_threshold": args.proposal_classifier_threshold if args.proposal_classifier else None,
|
"proposal_classifier_threshold": args.proposal_classifier_threshold if args.proposal_classifier else None,
|
||||||
"proposal_crop_scale": args.proposal_crop_scale if args.proposal_classifier else None,
|
"proposal_crop_scale": args.proposal_crop_scale if args.proposal_classifier else None,
|
||||||
|
"proposal_classifier_batch": args.proposal_classifier_batch if args.proposal_classifier else None,
|
||||||
"tile_count": len(tiles),
|
"tile_count": len(tiles),
|
||||||
"sweeps": sweeps,
|
"sweeps": sweeps,
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user