109 lines
2.5 KiB
Python
109 lines
2.5 KiB
Python
import torch
|
|
from PIL import Image
|
|
from transformers import Owlv2ForObjectDetection, Owlv2Processor
|
|
|
|
from app.config.settings import settings
|
|
from app.ml.image.weapon_prediction_result import WeaponPredictionResult
|
|
|
|
|
|
class WeaponDetector:
|
|
|
|
QUERIES = [
|
|
"a gun",
|
|
"a pistol",
|
|
"a revolver",
|
|
"a rifle",
|
|
"a shotgun",
|
|
"an assault rifle",
|
|
"a knife",
|
|
"a machete",
|
|
"a sword",
|
|
"a bomb",
|
|
"a grenade",
|
|
"a crossbow",
|
|
"a bow and arrow",
|
|
"brass knuckles",
|
|
"a nunchaku"
|
|
]
|
|
|
|
def __init__(self):
|
|
self.device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
print(f"Loading OWLv2 weapon detector on {self.device}")
|
|
|
|
self.model = Owlv2ForObjectDetection.from_pretrained(
|
|
settings.WEAPON_MODEL
|
|
)
|
|
self.processor = Owlv2Processor.from_pretrained(
|
|
settings.WEAPON_MODEL
|
|
)
|
|
|
|
self.model.to(self.device)
|
|
self.model.eval()
|
|
|
|
def predict(
|
|
self,
|
|
image: Image.Image
|
|
) -> WeaponPredictionResult:
|
|
inputs = self.processor(
|
|
text=self.QUERIES,
|
|
images=image,
|
|
return_tensors="pt"
|
|
)
|
|
|
|
inputs = {
|
|
key: value.to(self.device)
|
|
for key, value in inputs.items()
|
|
}
|
|
|
|
with torch.no_grad():
|
|
outputs = self.model(**inputs)
|
|
|
|
target_sizes = torch.tensor([image.size[::-1]])
|
|
|
|
results = self.processor.post_process_grounded_object_detection(
|
|
outputs,
|
|
threshold=0.0,
|
|
target_sizes=target_sizes
|
|
)[0]
|
|
|
|
scores = results["scores"].tolist()
|
|
labels = results["labels"].tolist()
|
|
boxes = results["boxes"].tolist()
|
|
|
|
image_area = image.width * image.height
|
|
label_scores = {}
|
|
|
|
for score, label_index, box in zip(scores, labels, boxes):
|
|
box_area = (box[2] - box[0]) * (box[3] - box[1])
|
|
|
|
if box_area < image_area * settings.WEAPON_MIN_AREA_FRACTION:
|
|
continue
|
|
|
|
if score >= settings.WEAPON_THRESHOLD:
|
|
query = self.QUERIES[label_index]
|
|
label_scores[query] = max(
|
|
label_scores.get(query, 0.0),
|
|
score
|
|
)
|
|
|
|
if not label_scores:
|
|
return WeaponPredictionResult(
|
|
label="normal",
|
|
score=0.0
|
|
)
|
|
|
|
best_label = max(
|
|
label_scores,
|
|
key=label_scores.get
|
|
)
|
|
|
|
print("====================")
|
|
print("WEAPON RESULT")
|
|
print("Scores:", label_scores)
|
|
|
|
return WeaponPredictionResult(
|
|
label=best_label,
|
|
score=label_scores[best_label],
|
|
detected_labels=sorted(label_scores.keys()),
|
|
box_count=len(label_scores)
|
|
) |