You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
343 lines
12 KiB
Python
343 lines
12 KiB
Python
import time
|
|
import threading
|
|
from collections import deque
|
|
|
|
import cv2
|
|
import numpy as np
|
|
import torch
|
|
|
|
from config import *
|
|
from helpers import clip_box, box_center, box_area
|
|
|
|
|
|
class AutoGazeROIWorker:
|
|
def __init__(self):
|
|
self.enabled = bool(AUTOGAZE_ENABLE)
|
|
self.ready = False
|
|
self.last_error = ""
|
|
|
|
self._autogaze_cls = None
|
|
self._processor_cls = None
|
|
if self.enabled:
|
|
try:
|
|
from autogaze.models.autogaze import AutoGaze, AutoGazeImageProcessor
|
|
self._autogaze_cls = AutoGaze
|
|
self._processor_cls = AutoGazeImageProcessor
|
|
except Exception as e:
|
|
self.last_error = f"import failed: {e}"
|
|
|
|
self.device = torch.device(f"cuda:{int(DEVICE)}") if torch.cuda.is_available() else torch.device("cpu")
|
|
|
|
self.clip_len = int(max(1, AUTOGAZE_CLIP_LEN))
|
|
self.update_every = int(max(1, AUTOGAZE_UPDATE_EVERY))
|
|
self.result_ttl = int(max(1, AUTOGAZE_RESULT_TTL))
|
|
self.topk = int(max(1, AUTOGAZE_TOPK))
|
|
self.min_active_cells = int(max(1, AUTOGAZE_MIN_ACTIVE_CELLS))
|
|
self.margin_cells = int(max(0, AUTOGAZE_ROI_MARGIN_CELLS))
|
|
self.min_side = float(max(8.0, AUTOGAZE_ROI_MIN_SIDE))
|
|
self.gazing_ratio = float(min(1.0, max(0.01, AUTOGAZE_GAZING_RATIO)))
|
|
self.task_loss_requirement = float(min(1.0, max(0.0, AUTOGAZE_TASK_LOSS_REQUIREMENT)))
|
|
|
|
self.model = None
|
|
self.processor = None
|
|
|
|
self.frame_buf = deque(maxlen=self.clip_len)
|
|
self.req = deque(maxlen=1)
|
|
self.lock = threading.Lock()
|
|
self.running = False
|
|
self.thread = threading.Thread(target=self._loop, daemon=True)
|
|
self.latest = None
|
|
self.last_submit_frame = -10**9
|
|
|
|
def status_line(self):
|
|
if not self.enabled:
|
|
return "AutoGaze disabled by config"
|
|
if self.ready:
|
|
return f"AutoGaze ready: model={AUTOGAZE_MODEL_ID} device={self.device}"
|
|
if self.last_error:
|
|
return f"AutoGaze unavailable: {self.last_error}"
|
|
return "AutoGaze unavailable"
|
|
|
|
def start(self):
|
|
if not self.enabled:
|
|
return
|
|
if self._autogaze_cls is None or self._processor_cls is None:
|
|
return
|
|
if not self._load_model():
|
|
return
|
|
if self.running:
|
|
return
|
|
self.running = True
|
|
self.thread.start()
|
|
|
|
def stop(self):
|
|
self.running = False
|
|
if self.thread.is_alive():
|
|
self.thread.join(timeout=1.0)
|
|
|
|
def submit(self, frame_eff_bgr, frame_id):
|
|
if (not self.running) or (not self.ready):
|
|
return
|
|
|
|
frame_copy = np.ascontiguousarray(frame_eff_bgr.copy())
|
|
fid = int(frame_id)
|
|
with self.lock:
|
|
self.frame_buf.append((fid, frame_copy))
|
|
if len(self.frame_buf) < self.clip_len:
|
|
return
|
|
if (fid - self.last_submit_frame) < self.update_every:
|
|
return
|
|
items = list(self.frame_buf)
|
|
clip = [f for _, f in items]
|
|
end_fid = int(items[-1][0])
|
|
self.req.append((clip, end_fid))
|
|
self.last_submit_frame = fid
|
|
|
|
def get_latest(self, current_frame_id, frame_w, frame_h, pred_ref_box=None):
|
|
info = {
|
|
"enabled": int(self.enabled),
|
|
"ready": int(self.ready),
|
|
"status": self.status_line(),
|
|
"stale": 1,
|
|
"age": -1,
|
|
"active_cells": 0,
|
|
"infer_ms": 0.0,
|
|
}
|
|
if not self.ready:
|
|
return None, [], info
|
|
|
|
with self.lock:
|
|
latest = None if self.latest is None else dict(self.latest)
|
|
|
|
if latest is None:
|
|
return None, [], info
|
|
|
|
age = int(current_frame_id) - int(latest["frame_id"])
|
|
info["age"] = age
|
|
info["active_cells"] = int(latest.get("active_cells", 0))
|
|
info["infer_ms"] = float(latest.get("infer_ms", 0.0))
|
|
if age > self.result_ttl:
|
|
return None, [], info
|
|
|
|
rois = []
|
|
for r in latest.get("rois", []):
|
|
rois.append(clip_box(np.array(r, dtype=np.float32), frame_w, frame_h))
|
|
|
|
roi = self._pick_roi(rois, pred_ref_box)
|
|
info["stale"] = 0
|
|
return roi, rois, info
|
|
|
|
def _load_model(self):
|
|
if self.ready and self.model is not None and self.processor is not None:
|
|
return True
|
|
|
|
try:
|
|
self.processor = self._processor_cls.from_pretrained(
|
|
AUTOGAZE_MODEL_ID,
|
|
local_files_only=bool(AUTOGAZE_LOCAL_FILES_ONLY),
|
|
)
|
|
self.model = self._autogaze_cls.from_pretrained(
|
|
AUTOGAZE_MODEL_ID,
|
|
use_flash_attn=bool(AUTOGAZE_USE_FLASH_ATTN),
|
|
local_files_only=bool(AUTOGAZE_LOCAL_FILES_ONLY),
|
|
)
|
|
self.model.to(self.device)
|
|
self.model.eval()
|
|
self.ready = True
|
|
self.last_error = ""
|
|
return True
|
|
except Exception as e:
|
|
self.model = None
|
|
self.processor = None
|
|
self.ready = False
|
|
self.last_error = f"load failed: {e}"
|
|
return False
|
|
|
|
def _loop(self):
|
|
while self.running:
|
|
item = None
|
|
with self.lock:
|
|
if self.req:
|
|
item = self.req.pop()
|
|
self.req.clear()
|
|
|
|
if item is None:
|
|
time.sleep(0.001)
|
|
continue
|
|
|
|
clip_frames, frame_id = item
|
|
infer_ms = 0.0
|
|
rois = []
|
|
active_cells = 0
|
|
try:
|
|
t0 = time.perf_counter()
|
|
rois, active_cells = self._infer_rois(clip_frames)
|
|
infer_ms = (time.perf_counter() - t0) * 1000.0
|
|
except Exception as e:
|
|
self.last_error = f"infer failed: {e}"
|
|
rois = []
|
|
active_cells = 0
|
|
infer_ms = 0.0
|
|
|
|
with self.lock:
|
|
self.latest = {
|
|
"frame_id": int(frame_id),
|
|
"rois": rois,
|
|
"active_cells": int(active_cells),
|
|
"infer_ms": float(infer_ms),
|
|
}
|
|
|
|
def _infer_rois(self, clip_frames_bgr):
|
|
if (self.model is None) or (self.processor is None):
|
|
return [], 0
|
|
if len(clip_frames_bgr) < self.clip_len:
|
|
return [], 0
|
|
|
|
h, w = clip_frames_bgr[-1].shape[:2]
|
|
clip_rgb = [cv2.cvtColor(f, cv2.COLOR_BGR2RGB) for f in clip_frames_bgr[-self.clip_len:]]
|
|
processed = self.processor([clip_rgb], return_tensors="pt")
|
|
pixel_values = processed.get("pixel_values", None)
|
|
if pixel_values is None:
|
|
return [], 0
|
|
if isinstance(pixel_values, list):
|
|
pixel_values = torch.as_tensor(np.asarray(pixel_values))
|
|
if pixel_values.ndim == 4:
|
|
pixel_values = pixel_values.unsqueeze(0)
|
|
|
|
video = pixel_values.to(self.device)
|
|
if video.dtype != torch.float32 and self.device.type == "cpu":
|
|
video = video.float()
|
|
|
|
with torch.inference_mode():
|
|
outputs = self.model(
|
|
{"video": video},
|
|
gazing_ratio=self.gazing_ratio,
|
|
task_loss_requirement=self.task_loss_requirement,
|
|
generate_only=True,
|
|
)
|
|
|
|
grid = self._extract_fine_grid(outputs)
|
|
if grid is None:
|
|
return [], 0
|
|
|
|
rois, active_cells = self._grid_to_rois(grid, w, h)
|
|
return rois, active_cells
|
|
|
|
def _extract_fine_grid(self, outputs):
|
|
gazing_mask = outputs.get("gazing_mask", None)
|
|
if gazing_mask is None or len(gazing_mask) == 0:
|
|
return None
|
|
|
|
fine = gazing_mask[-1]
|
|
if isinstance(fine, torch.Tensor):
|
|
arr = fine.detach().float().cpu().numpy()
|
|
else:
|
|
arr = np.asarray(fine)
|
|
|
|
if arr.ndim == 2:
|
|
arr = arr[None, ...]
|
|
if arr.ndim != 3 or arr.shape[0] <= 0 or arr.shape[1] <= 0:
|
|
return None
|
|
|
|
vec = arr[0, -1]
|
|
if vec.size <= 0:
|
|
return None
|
|
|
|
side = int(round(np.sqrt(float(vec.size))))
|
|
if side * side != int(vec.size):
|
|
return None
|
|
|
|
grid = vec.reshape(side, side)
|
|
return (grid > 0.5).astype(np.uint8)
|
|
|
|
def _grid_to_rois(self, grid, frame_w, frame_h):
|
|
if grid is None or grid.size == 0:
|
|
return [], 0
|
|
|
|
mask = (grid > 0).astype(np.uint8)
|
|
active_cells = int(mask.sum())
|
|
if active_cells < self.min_active_cells:
|
|
return [], active_cells
|
|
|
|
gh, gw = mask.shape
|
|
cell_w = float(frame_w) / float(max(1, gw))
|
|
cell_h = float(frame_h) / float(max(1, gh))
|
|
|
|
candidates = []
|
|
n_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask, connectivity=4)
|
|
for lbl in range(1, int(n_labels)):
|
|
x = int(stats[lbl, cv2.CC_STAT_LEFT])
|
|
y = int(stats[lbl, cv2.CC_STAT_TOP])
|
|
ww = int(stats[lbl, cv2.CC_STAT_WIDTH])
|
|
hh = int(stats[lbl, cv2.CC_STAT_HEIGHT])
|
|
area = int(stats[lbl, cv2.CC_STAT_AREA])
|
|
if area < self.min_active_cells:
|
|
continue
|
|
|
|
gx1 = max(0, x - self.margin_cells)
|
|
gy1 = max(0, y - self.margin_cells)
|
|
gx2 = min(gw, x + ww + self.margin_cells)
|
|
gy2 = min(gh, y + hh + self.margin_cells)
|
|
roi = np.array(
|
|
[gx1 * cell_w, gy1 * cell_h, gx2 * cell_w, gy2 * cell_h],
|
|
dtype=np.float32,
|
|
)
|
|
roi = self._ensure_min_side(roi, frame_w, frame_h)
|
|
candidates.append((area, roi))
|
|
|
|
if not candidates:
|
|
ys, xs = np.where(mask > 0)
|
|
if xs.size <= 0:
|
|
return [], active_cells
|
|
gx1 = max(0, int(xs.min()) - self.margin_cells)
|
|
gy1 = max(0, int(ys.min()) - self.margin_cells)
|
|
gx2 = min(gw, int(xs.max()) + 1 + self.margin_cells)
|
|
gy2 = min(gh, int(ys.max()) + 1 + self.margin_cells)
|
|
roi = np.array(
|
|
[gx1 * cell_w, gy1 * cell_h, gx2 * cell_w, gy2 * cell_h],
|
|
dtype=np.float32,
|
|
)
|
|
roi = self._ensure_min_side(roi, frame_w, frame_h)
|
|
candidates.append((active_cells, roi))
|
|
|
|
candidates.sort(key=lambda x: x[0], reverse=True)
|
|
rois = [np.array(r, dtype=np.float32) for _, r in candidates[: self.topk]]
|
|
return rois, active_cells
|
|
|
|
def _ensure_min_side(self, roi, frame_w, frame_h):
|
|
roi = clip_box(roi, frame_w, frame_h)
|
|
rw = float(roi[2] - roi[0])
|
|
rh = float(roi[3] - roi[1])
|
|
if rw >= self.min_side and rh >= self.min_side:
|
|
return roi
|
|
|
|
cx = 0.5 * (float(roi[0]) + float(roi[2]))
|
|
cy = 0.5 * (float(roi[1]) + float(roi[3]))
|
|
side = max(self.min_side, rw, rh)
|
|
return clip_box(
|
|
np.array([cx - 0.5 * side, cy - 0.5 * side, cx + 0.5 * side, cy + 0.5 * side], dtype=np.float32),
|
|
frame_w,
|
|
frame_h,
|
|
)
|
|
|
|
def _pick_roi(self, rois, pred_ref_box):
|
|
if not rois:
|
|
return None
|
|
if pred_ref_box is None:
|
|
return rois[0]
|
|
|
|
pref = np.array(pred_ref_box, dtype=np.float32)
|
|
pc = box_center(pref)
|
|
|
|
best = None
|
|
best_key = (1e9, 1e9)
|
|
for r in rois:
|
|
c = box_center(r)
|
|
dist = float(np.linalg.norm(c - pc))
|
|
a = float(box_area(r))
|
|
key = (dist, -a)
|
|
if key < best_key:
|
|
best_key = key
|
|
best = r
|
|
return best
|