feat: live count worker, tracking, and page updates

This commit is contained in:
Andrew-AAAA committed 2026-09-30 11:31:13 +07:00
1 parent 4606c0adeb
commit 1403a07b86
5 files changed
+77 -15

No files matched your search

+5 -3
View File
@@ -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()
+2
View File
@@ -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))
+5 -2
View File
@@ -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
+9 -5
View File
@@ -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()
+56 -5
View File
@@ -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'}