feat: live count worker, tracking, and page updates
This commit is contained in:
1 parent
4606c0adeb
commit
1403a07b86
5 files changed
+77
-15
No files matched your search
@@ -28,7 +28,8 @@ _TRACKER_CFG = os.path.join(
|
||||
class ByteTrackTracker:
|
||||
"""Tracks sacks across frames using FastTrack (occlusion-aware)."""
|
||||
|
||||
def __init__(self, model_path_or_model: str | YOLO, conf: float = 0.35) -> None:
|
||||
def __init__(self, model_path_or_model: str | YOLO, conf: float = 0.35,
|
||||
class_filter: tuple = ("sack", "truck")) -> None:
|
||||
if isinstance(model_path_or_model, str):
|
||||
self._model = YOLO(model_path_or_model)
|
||||
self._model_path = model_path_or_model
|
||||
@@ -36,6 +37,7 @@ class ByteTrackTracker:
|
||||
self._model = model_path_or_model
|
||||
self._model_path = model_path_or_model.ckpt_path if hasattr(model_path_or_model, 'ckpt_path') else ""
|
||||
self._conf = conf
|
||||
self._class_filter = class_filter
|
||||
self._tracker_cfg = _TRACKER_CFG
|
||||
|
||||
def update(
|
||||
@@ -57,8 +59,8 @@ class ByteTrackTracker:
|
||||
for i, box in enumerate(result.boxes):
|
||||
cls_id = int(box.cls[0])
|
||||
name = self._model.names[cls_id]
|
||||
# Retain both 'sack' and 'truck' classes
|
||||
if name not in ("sack", "truck"):
|
||||
# Retain only the classes the caller asked to track
|
||||
if name not in self._class_filter:
|
||||
continue
|
||||
track_id = int(ids[i]) if ids is not None else None
|
||||
x1, y1, x2, y2 = box.xyxy[0].tolist()
|
||||
|
||||
@@ -22,6 +22,7 @@ class StartRequest(BaseModel):
|
||||
source_rel: Optional[str] = None
|
||||
model_path: Optional[str] = None
|
||||
model_version_id: Optional[int] = None
|
||||
count_classes: list = ["sack"]
|
||||
line_y: int = 266
|
||||
line_x_start: int = 469
|
||||
line_x_end: int = 910
|
||||
@@ -97,6 +98,7 @@ def start(project_id: int, request: StartRequest) -> dict:
|
||||
min_area_scale=request.min_area_scale,
|
||||
spatial_dedup=request.spatial_dedup,
|
||||
whep=whep,
|
||||
count_classes=tuple(request.count_classes),
|
||||
)
|
||||
except live_count.LiveCountError as exc:
|
||||
raise HTTPException(400, str(exc))
|
||||
|
||||
@@ -25,6 +25,7 @@ DEFAULTS = {
|
||||
"entry_travel_min": 60.0, "handoff_radius": 100.0,
|
||||
"unload_confirm_frames": 3, "min_area_scale": 1.0,
|
||||
"dedup_radius": 60.0, "spatial_dedup": False,
|
||||
"count_classes": ["sack"],
|
||||
}
|
||||
|
||||
|
||||
@@ -176,7 +177,8 @@ def count_video(path: str, model, params: dict, should_cancel=None,
|
||||
raise CountingBenchError(f"Could not open {path}")
|
||||
total_frames = int(capture.get(cv2.CAP_PROP_FRAME_COUNT) or 0)
|
||||
|
||||
tracker = ByteTrackTracker(model, settings["conf"])
|
||||
tracker = ByteTrackTracker(model, settings["conf"],
|
||||
class_filter=tuple(settings["count_classes"]))
|
||||
stabilizer = BboxStabilizer(ema_alpha=0.35, max_hold_frames=10,
|
||||
max_height_ratio=1.5, min_height_ratio=0.70)
|
||||
counter = LineCrossCounter(
|
||||
@@ -199,7 +201,8 @@ def count_video(path: str, model, params: dict, should_cancel=None,
|
||||
if not ok or frame is None:
|
||||
break
|
||||
frame = cv2.resize(frame, (1280, 720))
|
||||
detections = [d for d in tracker.update(frame, []) if d.class_name == "sack"]
|
||||
# The tracker only emits count_classes, so no post-filter is needed.
|
||||
detections = tracker.update(frame, [])
|
||||
stable = stabilizer.update(detections)
|
||||
inside = [
|
||||
d for d in stable
|
||||
|
||||
@@ -43,7 +43,8 @@ class Session:
|
||||
dedup_radius: float, margin: int, imgsz: int,
|
||||
entry_travel_min: float, handoff_radius: float,
|
||||
unload_confirm_frames: int, min_area_scale: float,
|
||||
spatial_dedup: bool, whep_url: str = ""):
|
||||
spatial_dedup: bool, whep_url: str = "",
|
||||
count_classes: tuple = ("sack",)):
|
||||
self.source = source
|
||||
# When the browser watches the camera over WebRTC it never asks for the
|
||||
# MJPEG, so encoding a JPEG per frame would be pure waste (REQ-177).
|
||||
@@ -61,6 +62,7 @@ class Session:
|
||||
self.unload_confirm_frames = unload_confirm_frames
|
||||
self.min_area_scale = min_area_scale
|
||||
self.spatial_dedup = spatial_dedup
|
||||
self.count_classes = tuple(count_classes)
|
||||
|
||||
self.started_at = time.time()
|
||||
self.error = ""
|
||||
@@ -167,7 +169,8 @@ class Session:
|
||||
# segfault without it (same reason predict.py does this).
|
||||
model(np.zeros((720, 1280, 3), dtype=np.uint8), imgsz=self.imgsz, verbose=False)
|
||||
|
||||
tracker = ByteTrackTracker(model, self.conf)
|
||||
tracker = ByteTrackTracker(model, self.conf,
|
||||
class_filter=tuple(self.count_classes))
|
||||
stabilizer = BboxStabilizer(ema_alpha=0.35, max_hold_frames=10,
|
||||
max_height_ratio=1.5, min_height_ratio=0.70)
|
||||
counter = LineCrossCounter(
|
||||
@@ -200,7 +203,8 @@ class Session:
|
||||
break
|
||||
|
||||
frame = cv2.resize(frame, (1280, 720))
|
||||
detections = [d for d in tracker.update(frame, []) if d.class_name == "sack"]
|
||||
# The tracker only emits count_classes, so no post-filter is needed.
|
||||
detections = tracker.update(frame, [])
|
||||
stable = stabilizer.update(detections)
|
||||
inside, outside = [], []
|
||||
small = 0
|
||||
@@ -338,7 +342,7 @@ def start(source: str, model_path: str, line_y: int, line_x_start: int, line_x_e
|
||||
imgsz: int = 640, entry_travel_min: float = 60.0,
|
||||
handoff_radius: float = 100.0, unload_confirm_frames: int = 3,
|
||||
min_area_scale: float = 1.0, spatial_dedup: bool = False,
|
||||
whep: str = "") -> dict:
|
||||
whep: str = "", count_classes: tuple = ("sack",)) -> dict:
|
||||
global _session
|
||||
with _guard:
|
||||
if _session is not None and _session.status()["running"]:
|
||||
@@ -350,7 +354,7 @@ def start(source: str, model_path: str, line_y: int, line_x_start: int, line_x_e
|
||||
_session = Session(source, model_path, line_y, line_x_start, line_x_end,
|
||||
conf, dedup_radius, margin, imgsz, entry_travel_min,
|
||||
handoff_radius, unload_confirm_frames, min_area_scale,
|
||||
spatial_dedup, whep)
|
||||
spatial_dedup, whep, count_classes)
|
||||
_session.start()
|
||||
time.sleep(0.4) # let an immediate failure surface in the response
|
||||
return _session.status()
|
||||
|
||||
@@ -32,12 +32,14 @@ export default function LiveCountPage({ projectId, onProject }) {
|
||||
const [videoRel, setVideoRel] = useState('')
|
||||
// Defaults are the settings that were dialled in against the real camera —
|
||||
// a fresh session starts where the last tuning session left off.
|
||||
// count_classes is the multi-select of which detected classes actually count.
|
||||
const [cfg, setCfg] = useState({
|
||||
line_y: 266, line_x_start: 469, line_x_end: 910,
|
||||
margin: 5, dedup_radius: 60, conf: 0.35,
|
||||
entry_travel_min: 60, handoff_radius: 100, unload_confirm_frames: 3,
|
||||
min_area_scale: 1.0,
|
||||
min_area_scale: 1.0, count_classes: ['sack'],
|
||||
})
|
||||
const [projectClasses, setProjectClasses] = useState([])
|
||||
const [status, setStatus] = useState({ running: false })
|
||||
const [error, setError] = useState('')
|
||||
const [busy, setBusy] = useState(false)
|
||||
@@ -46,7 +48,10 @@ export default function LiveCountPage({ projectId, onProject }) {
|
||||
const pollRef = useRef(null)
|
||||
|
||||
useEffect(() => {
|
||||
api.getProject(projectId).then((p) => onProject?.(p)).catch(() => {})
|
||||
api.getProject(projectId).then((p) => {
|
||||
setProjectClasses(p.classes || [])
|
||||
onProject?.(p)
|
||||
}).catch(() => {})
|
||||
api.liveCountModels(projectId)
|
||||
.then((payload) => {
|
||||
setModels(payload.models)
|
||||
@@ -121,10 +126,21 @@ export default function LiveCountPage({ projectId, onProject }) {
|
||||
}
|
||||
}
|
||||
|
||||
function toggleCountClass(name) {
|
||||
setCfg((prev) => {
|
||||
const current = prev.count_classes || []
|
||||
return {
|
||||
...prev,
|
||||
count_classes: current.includes(name)
|
||||
? current.filter((c) => c !== name)
|
||||
: [...current, name],
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Click on the stream to place whichever edge is armed. Placing by eye beats
|
||||
// guessing a pixel value on a slider.
|
||||
function placeOnClick(event) {
|
||||
if (!running) return
|
||||
function placeOnClick(event) { if (!running) return
|
||||
const rect = event.currentTarget.getBoundingClientRect()
|
||||
if (!rect.height || !rect.width) return
|
||||
if (placing === 'line_y') {
|
||||
@@ -253,6 +269,40 @@ export default function LiveCountPage({ projectId, onProject }) {
|
||||
</select>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<label className="hint" style={{ fontSize: '0.8rem' }}>Count classes</label>
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 6, marginTop: 6 }}>
|
||||
{projectClasses.length === 0 && (
|
||||
<span className="hint" style={{ fontSize: '0.76rem' }}>no classes on this project</span>
|
||||
)}
|
||||
{projectClasses.map((cls) => {
|
||||
const on = (cfg.count_classes || []).includes(cls.name)
|
||||
return (
|
||||
<button
|
||||
key={cls.class_id}
|
||||
type="button"
|
||||
disabled={running}
|
||||
onClick={() => toggleCountClass(cls.name)}
|
||||
style={{
|
||||
fontSize: '0.76rem', padding: '3px 10px', borderRadius: 6,
|
||||
cursor: running ? 'not-allowed' : 'pointer',
|
||||
background: on ? 'rgba(192, 132, 252, 0.2)' : 'rgba(255,255,255,0.05)',
|
||||
color: on ? '#f3e8ff' : '#a1a1aa',
|
||||
border: on ? '1px solid rgba(192, 132, 252, 0.5)' : '1px solid #3f3f46',
|
||||
}}
|
||||
>
|
||||
{cls.name}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
{(cfg.count_classes || []).length === 0 && (
|
||||
<p className="hint" style={{ color: '#ef4444', fontSize: '0.74rem', margin: '4px 0 0' }}>
|
||||
Select at least one class to count.
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{FIELDS.map((f) => {
|
||||
const locked = running && !f.live
|
||||
return (
|
||||
@@ -287,7 +337,8 @@ export default function LiveCountPage({ projectId, onProject }) {
|
||||
<button
|
||||
className="btn btn-primary"
|
||||
onClick={start}
|
||||
disabled={busy || !modelPath || (mode === 'file' ? !videoRel : !source)}
|
||||
disabled={busy || !modelPath || (cfg.count_classes || []).length === 0
|
||||
|| (mode === 'file' ? !videoRel : !source)}
|
||||
style={{ flex: 1, cursor: busy || !modelPath ? 'not-allowed' : 'pointer' }}
|
||||
>
|
||||
{busy ? 'Starting…' : 'Start counting'}
|
||||
|
||||
Reference in new issue
Block a user