commit a2d391d38bce9ffa2f498a63aafd470f6b3ff45f
parent bbe60dbdcccf08da32c6a2933872c6d6aa7fb3e3
Author: david cochran <about.trout@gmail.com>
Date: Wed, 14 Feb 2024 10:55:15 -0800
add GPU support for boatfinder
Diffstat:
2 files changed, 9 insertions(+), 5 deletions(-)
diff --git a/README.md b/README.md
@@ -7,7 +7,8 @@ Locate large boats as they pass by Pillar Point.
Install the requirements.
```
-$ sudo apt install ffmpeg
+$ sudo apt install ffmpeg # for extracing keyframes
+$ sudo apt install nvidia-cuda-toolkit # for GPU support
$ python3 -m venv .venv
$ source .venv/bin/activate
$ pip install -r requirements.txt
diff --git a/boatfinder.py b/boatfinder.py
@@ -14,15 +14,17 @@ class BoatFinder:
logging.info(f"BoatFinder initializing ...")
t0 = perf_counter()
# https://huggingface.co/facebook/detr-resnet-50
+ self.device = "cuda:0" if torch.cuda.is_available() else "cpu"
self.processor = DetrImageProcessor.from_pretrained(
"facebook/detr-resnet-50", revision="no_timm"
)
self.model = DetrForObjectDetection.from_pretrained(
"facebook/detr-resnet-50", revision="no_timm"
- )
- logging.info(f"BoatFinder initialized! took {perf_counter() - t0} seconds")
+ ).to(self.device)
+ logging.info(f"BoatFinder initialized (for {self.device}); took {perf_counter() - t0} seconds")
def find(self, img):
+ t0 = perf_counter()
res = self.__match_results(img)
for score, label, box in zip(res["scores"], res["labels"], res["boxes"]):
label = self.model.config.id2label[label.item()]
@@ -30,11 +32,12 @@ class BoatFinder:
score = round(score.item(), 3)
box = [round(i, 3) for i in box.tolist()]
yield (score, label, box)
+ logging.info(f"Finished search in {perf_counter() - t0} seconds")
def __match_results(self, image):
- inputs = self.processor(images=image, return_tensors="pt")
+ inputs = self.processor(images=image, return_tensors="pt").to(self.device)
outputs = self.model(**inputs)
- target_sizes = torch.tensor([image.size[::-1]])
+ target_sizes = torch.tensor([image.size[::-1]]).to(self.device)
return self.processor.post_process_object_detection(
outputs, target_sizes=target_sizes, threshold=0.5
)[0]