polling_tracker.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459
  1. """
  2. 轮询跟踪 + PTZ 抓拍协调器
  3. """
  4. import os
  5. import time
  6. import threading
  7. import logging
  8. from typing import Dict, List, Tuple, Optional
  9. from dataclasses import dataclass
  10. import cv2
  11. import numpy as np
  12. from config import TRACKING_CONFIG
  13. from tracker import UltralyticsTracker, TrackedPerson
  14. from coordinator import TargetSelector, TrackingTarget
  15. logger = logging.getLogger(__name__)
  16. @dataclass
  17. class CaptureRecord:
  18. """单次抓拍记录"""
  19. track_id: int
  20. timestamp: float
  21. position: Tuple[float, float]
  22. ptz_position: Tuple[float, float, int]
  23. ptz_image: np.ndarray
  24. panorama_image: Optional[np.ndarray]
  25. confidence: float
  26. class PollingTrackingCoordinator:
  27. """多目标轮询跟踪 + PTZ 抓拍协调器"""
  28. def __init__(
  29. self,
  30. panorama_camera,
  31. ptz_camera,
  32. tracker: UltralyticsTracker,
  33. config: Optional[Dict] = None,
  34. calibrator=None,
  35. ):
  36. self.panorama = panorama_camera
  37. self.ptz = ptz_camera
  38. self.tracker = tracker
  39. self.config = config or TRACKING_CONFIG
  40. self.calibrator = calibrator
  41. self.active_targets: Dict[int, TrackedPerson] = {}
  42. self.target_order: List[int] = []
  43. self.current_index: int = 0
  44. self.batch_captures: List[CaptureRecord] = []
  45. self._capture_counts: Dict[int, int] = {}
  46. self._last_seen_time: Dict[int, float] = {}
  47. self.targets_lock = threading.Lock()
  48. self.batch_lock = threading.Lock()
  49. self._capture_counts_lock = threading.Lock()
  50. self.running = False
  51. self._detection_thread = None
  52. self._ptz_thread = None
  53. self._paused = False
  54. self._paused_event = threading.Event()
  55. self._paused_event.set()
  56. self._last_ptz_command_time = 0.0
  57. self._last_ptz_command_time_lock = threading.Lock()
  58. self.target_selector = TargetSelector(self.config.get("target_selection", {}))
  59. self.event_pusher = None
  60. self.stats = {
  61. "frames_processed": 0,
  62. "persons_detected": 0,
  63. "captures": 0,
  64. "uploads": 0,
  65. "start_time": None,
  66. }
  67. self.stats_lock = threading.Lock()
  68. self._ensure_capture_dir()
  69. def set_event_pusher(self, event_pusher):
  70. self.event_pusher = event_pusher
  71. def _ensure_capture_dir(self):
  72. capture_dir = self.config.get("capture_dir", "/home/admin/dsh/tracking_captures")
  73. try:
  74. os.makedirs(capture_dir, exist_ok=True)
  75. except OSError as e:
  76. logger.warning(f"无法创建抓拍目录 {capture_dir}: {e}")
  77. def start(self) -> bool:
  78. if not self.panorama.connect():
  79. logger.error("连接全景摄像头失败")
  80. return False
  81. if not self.ptz.connect():
  82. logger.error("连接球机失败")
  83. self.panorama.disconnect()
  84. return False
  85. if not self.panorama.start_stream_rtsp():
  86. logger.error("启动全景视频流失败")
  87. self.panorama.disconnect()
  88. self.ptz.disconnect()
  89. return False
  90. self.running = True
  91. self._detection_thread = threading.Thread(target=self._detection_worker, daemon=True)
  92. self._detection_thread.start()
  93. self._ptz_thread = threading.Thread(target=self._ptz_worker, daemon=True)
  94. self._ptz_thread.start()
  95. with self.stats_lock:
  96. self.stats["start_time"] = time.time()
  97. logger.info("轮询跟踪抓拍协调器已启动")
  98. return True
  99. def stop(self):
  100. self.running = False
  101. self._paused_event.set()
  102. if self._detection_thread:
  103. self._detection_thread.join(timeout=3)
  104. if self._ptz_thread:
  105. self._ptz_thread.join(timeout=3)
  106. # 刷新待上传批次
  107. self._flush_batch_if_needed()
  108. # 停止视频流后再断开连接
  109. if hasattr(self.panorama, "stop_stream_rtsp"):
  110. try:
  111. self.panorama.stop_stream_rtsp()
  112. except Exception as e:
  113. logger.warning(f"停止全景视频流失败: {e}")
  114. self.panorama.disconnect()
  115. self.ptz.disconnect()
  116. # 释放跟踪器资源
  117. if self.tracker is not None and hasattr(self.tracker, "release"):
  118. try:
  119. self.tracker.release()
  120. except Exception as e:
  121. logger.warning(f"释放跟踪器失败: {e}")
  122. logger.info("轮询跟踪抓拍协调器已停止")
  123. def pause(self):
  124. self._paused = True
  125. self._paused_event.clear()
  126. def resume(self):
  127. self._paused = False
  128. self._paused_event.set()
  129. def _detection_worker(self):
  130. self._paused_event.wait()
  131. detection_fps = self.config.get("detection_fps", 2)
  132. detection_interval = 1.0 / detection_fps
  133. last_detection_time = 0
  134. while self.running:
  135. try:
  136. # 暂停时阻塞等待,避免忙等
  137. if self._paused:
  138. self._paused_event.wait()
  139. continue
  140. frame = self.panorama.get_frame()
  141. if frame is None:
  142. time.sleep(0.01)
  143. continue
  144. self._update_stats("frames_processed")
  145. current_time = time.time()
  146. if current_time - last_detection_time >= detection_interval:
  147. last_detection_time = current_time
  148. tracked = self.tracker.update(frame)
  149. self._update_active_targets(tracked, frame.shape)
  150. if tracked:
  151. self._update_stats("persons_detected", len(tracked))
  152. time.sleep(0.01)
  153. except Exception as e:
  154. logger.error(f"检测线程错误: {e}")
  155. time.sleep(0.1)
  156. def _update_active_targets(self, tracked: List[TrackedPerson], frame_shape):
  157. current_time = time.time()
  158. frame_h, frame_w = frame_shape[:2]
  159. timeout = self.config.get("tracking_timeout", 3.0)
  160. max_targets = self.config.get("max_tracking_targets", 4)
  161. with self.targets_lock:
  162. # 更新或新增
  163. updated_ids = set()
  164. for p in tracked:
  165. if p.track_id < 0:
  166. continue
  167. updated_ids.add(p.track_id)
  168. p.lost = False
  169. self.active_targets[p.track_id] = p
  170. self._last_seen_time[p.track_id] = current_time
  171. if p.track_id not in self.target_order:
  172. self.target_order.append(p.track_id)
  173. # 标记丢失
  174. for tid in self.target_order:
  175. if tid not in updated_ids:
  176. t = self.active_targets.get(tid)
  177. if t is not None:
  178. t.lost = True
  179. # 移除长期丢失(超过 tracking_timeout)
  180. remove_ids = []
  181. for tid in self.target_order:
  182. t = self.active_targets.get(tid)
  183. if t is None:
  184. remove_ids.append(tid)
  185. continue
  186. if t.lost:
  187. last_seen = self._last_seen_time.get(tid, current_time)
  188. if current_time - last_seen >= timeout:
  189. remove_ids.append(tid)
  190. for tid in remove_ids:
  191. if tid in self.target_order:
  192. self.target_order.remove(tid)
  193. self.active_targets.pop(tid, None)
  194. self._last_seen_time.pop(tid, None)
  195. with self._capture_counts_lock:
  196. self._capture_counts.pop(tid, None)
  197. # 人数上限淘汰
  198. if len(self.active_targets) > max_targets:
  199. self._prune_targets(frame_w, frame_h, max_targets)
  200. def _prune_targets(self, frame_w: int, frame_h: int, max_targets: int):
  201. targets = list(self.active_targets.values())
  202. frame_size = (frame_w, frame_h)
  203. scored = []
  204. for t in targets:
  205. area = (t.bbox[2] - t.bbox[0]) * (t.bbox[3] - t.bbox[1])
  206. center_distance = self._center_distance(t.center, frame_size)
  207. target = TrackingTarget(
  208. track_id=t.track_id,
  209. position=(t.center[0] / frame_w, t.center[1] / frame_h),
  210. last_update=time.time(),
  211. area=area,
  212. confidence=t.confidence,
  213. center_distance=center_distance,
  214. )
  215. target.score = self.target_selector.calculate_score(target, frame_size)
  216. scored.append(target)
  217. scored.sort(key=lambda x: x.score, reverse=True)
  218. keep_ids = {t.track_id for t in scored[:max_targets]}
  219. remove_ids = [tid for tid in self.active_targets if tid not in keep_ids]
  220. for tid in remove_ids:
  221. self.active_targets.pop(tid, None)
  222. self._last_seen_time.pop(tid, None)
  223. if tid in self.target_order:
  224. self.target_order.remove(tid)
  225. with self._capture_counts_lock:
  226. self._capture_counts.pop(tid, None)
  227. def _center_distance(self, center: Tuple[int, int], frame_size: Tuple[int, int]) -> float:
  228. cx, cy = frame_size[0] / 2, frame_size[1] / 2
  229. dx = abs(center[0] - cx) / cx
  230. dy = abs(center[1] - cy) / cy
  231. return (dx + dy) / 2
  232. def _ptz_worker(self):
  233. while self.running:
  234. try:
  235. # 暂停时阻塞等待恢复
  236. if self._paused:
  237. self._paused_event.wait()
  238. continue
  239. # 原子性获取目标快照和当前目标
  240. with self.targets_lock:
  241. target_order_snapshot = self.target_order.copy()
  242. has_targets = bool(self.active_targets)
  243. if target_order_snapshot:
  244. if self.current_index >= len(target_order_snapshot):
  245. self.current_index = 0
  246. target_id = target_order_snapshot[self.current_index]
  247. target = self.active_targets.get(target_id)
  248. else:
  249. target_id = None
  250. target = None
  251. if not has_targets or not target_order_snapshot:
  252. self._flush_batch_if_needed()
  253. time.sleep(0.1)
  254. continue
  255. if target is None or target.lost:
  256. # 目标丢失时跳过但保留在队列中,并短暂休眠避免忙等
  257. time.sleep(0.01)
  258. with self.targets_lock:
  259. self._advance(len(target_order_snapshot))
  260. continue
  261. record = self._capture_one(target)
  262. if record:
  263. with self.batch_lock:
  264. self.batch_captures.append(record)
  265. self._update_stats("captures")
  266. with self.targets_lock:
  267. self._advance(len(target_order_snapshot))
  268. # 一轮完成
  269. with self.batch_lock:
  270. if self.current_index == 0 and self.batch_captures:
  271. self._upload_batch(self.batch_captures)
  272. self.batch_captures.clear()
  273. except Exception as e:
  274. logger.error(f"PTZ 线程错误: {e}")
  275. time.sleep(0.1)
  276. def _advance(self, order_len: int = None):
  277. if order_len is None:
  278. order_len = len(self.target_order) or 1
  279. self.current_index = (self.current_index + 1) % order_len
  280. def _capture_one(self, target: TrackedPerson) -> Optional[CaptureRecord]:
  281. frame = self.panorama.get_frame()
  282. if frame is None:
  283. return None
  284. frame_h, frame_w = frame.shape[:2]
  285. x_ratio = target.center[0] / frame_w
  286. y_ratio = target.center[1] / frame_h
  287. if self.calibrator and self.calibrator.is_calibrated():
  288. pan, tilt = self.calibrator.transform(x_ratio, y_ratio)
  289. ptz_config = getattr(self.ptz, "ptz_config", {})
  290. if ptz_config.get("pan_flip"):
  291. pan = (pan + 180) % 360
  292. zoom = ptz_config.get("default_zoom", 8)
  293. else:
  294. pan, tilt, zoom = self.ptz.calculate_ptz_position(x_ratio, y_ratio)
  295. # enforce PTZ command cooldown
  296. ptz_command_cooldown = self.config.get("ptz_command_cooldown", 0.2)
  297. with self._last_ptz_command_time_lock:
  298. elapsed = time.time() - self._last_ptz_command_time
  299. if elapsed < ptz_command_cooldown:
  300. time.sleep(ptz_command_cooldown - elapsed)
  301. success = self.ptz.goto_exact_position(pan, tilt, zoom)
  302. if not success:
  303. return None
  304. with self._last_ptz_command_time_lock:
  305. self._last_ptz_command_time = time.time()
  306. ptz_stabilize_time = self.config.get("ptz_stabilize_time", 2.0)
  307. time.sleep(max(ptz_stabilize_time, ptz_command_cooldown))
  308. ptz_frame = self._get_clear_ptz_frame()
  309. if ptz_frame is None:
  310. return None
  311. max_cap = self.config.get("max_capture_per_target", 0)
  312. with self._capture_counts_lock:
  313. if max_cap > 0 and self._capture_counts.get(target.track_id, 0) >= max_cap:
  314. return None
  315. self._capture_counts[target.track_id] = self._capture_counts.get(target.track_id, 0) + 1
  316. panorama_image = frame.copy() if self.config.get("save_panorama_pair", True) else None
  317. # 本地保存
  318. self._save_local(ptz_frame, panorama_image, target, pan, tilt, zoom)
  319. return CaptureRecord(
  320. track_id=target.track_id,
  321. timestamp=time.time(),
  322. position=(x_ratio, y_ratio),
  323. ptz_position=(pan, tilt, zoom),
  324. ptz_image=ptz_frame,
  325. panorama_image=panorama_image,
  326. confidence=target.confidence,
  327. )
  328. def _get_clear_ptz_frame(self, max_attempts: int = 5, wait_interval: float = 0.2) -> Optional[np.ndarray]:
  329. best_frame = None
  330. best_score = -1
  331. for _ in range(max_attempts):
  332. frame = self.ptz.get_frame()
  333. if frame is not None:
  334. frame_copy = frame.copy()
  335. gray = cv2.cvtColor(frame_copy, cv2.COLOR_BGR2GRAY)
  336. score = cv2.Laplacian(gray, cv2.CV_64F).var()
  337. if score > best_score:
  338. best_score = score
  339. best_frame = frame_copy
  340. time.sleep(wait_interval)
  341. return best_frame
  342. def _save_local(self, ptz_frame, panorama_image, target, pan, tilt, zoom):
  343. capture_dir = self.config.get("capture_dir", "/home/admin/dsh/tracking_captures")
  344. try:
  345. os.makedirs(capture_dir, exist_ok=True)
  346. except OSError as e:
  347. logger.warning(f"无法创建抓拍目录 {capture_dir}: {e}")
  348. return
  349. timestamp = int(time.time() * 1000)
  350. base = f"{capture_dir}/ptz_{target.track_id}_{timestamp}_{pan:.0f}_{tilt:.0f}_z{zoom}.jpg"
  351. cv2.imwrite(base, ptz_frame)
  352. if panorama_image is not None:
  353. pan_base = f"{capture_dir}/panorama_{target.track_id}_{timestamp}.jpg"
  354. cv2.imwrite(pan_base, panorama_image)
  355. def _upload_batch(self, records: List[CaptureRecord]):
  356. if not self.event_pusher or not self.config.get("enable_upload", True):
  357. return
  358. try:
  359. uploads = []
  360. for r in records:
  361. ptz_url = self.event_pusher.upload_numpy_image(r.ptz_image)
  362. pan_url = None
  363. if r.panorama_image is not None:
  364. pan_url = self.event_pusher.upload_numpy_image(r.panorama_image)
  365. uploads.append({
  366. "track_id": r.track_id,
  367. "ptz_image_url": ptz_url,
  368. "panorama_image_url": pan_url,
  369. "position": r.position,
  370. "ptz_position": r.ptz_position,
  371. "confidence": r.confidence,
  372. "timestamp": r.timestamp,
  373. })
  374. self.event_pusher.push_tracking_capture(batch_time=time.time(), captures=uploads)
  375. self._update_stats("uploads")
  376. except Exception as e:
  377. logger.error(f"批量上传失败: {e}")
  378. def _flush_batch_if_needed(self):
  379. with self.batch_lock:
  380. if self.batch_captures:
  381. self._upload_batch(self.batch_captures)
  382. self.batch_captures.clear()
  383. def _update_stats(self, key: str, value: int = 1):
  384. with self.stats_lock:
  385. if key in self.stats:
  386. self.stats[key] += value
  387. def get_stats(self) -> Dict:
  388. with self.stats_lock:
  389. return self.stats.copy()