fix: exclude truck class from counting, use MultiClassLineCounter, real class labels
This commit is contained in:
1 parent
35cccd03b3
commit
d38375986e
4 files changed
+245
-6
No files matched your search
@@ -0,0 +1,155 @@
|
||||
# Truck Counting Bug Fix + Box Counts + Download Compression Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Fix truck-classed-as-sack counting bug, propagate sack vs box counts to results UI, compress downloads to max 250MB.
|
||||
|
||||
**Architecture:** (1) Filter truck class before counting and switch both pipelines to `MultiClassLineCounter` (defense in depth + box counts); fix dashboard hardcoded "sack" label. (2) Add box count fields to result dataclasses and UI. (3) Compress download on-the-fly via ffmpeg targeting 250MB.
|
||||
|
||||
**Tech Stack:** Python 3.10, Flask, OpenCV, ffmpeg subprocess, existing `MultiClassLineCounter`
|
||||
|
||||
**Spec:** Original repo reference: `/tmp/karung-counting-feedmill-semarang/predict.py` lines 1497-1500 (class filter), 1409 (`MultiClassLineCounter`)
|
||||
|
||||
## Global Constraints
|
||||
|
||||
- Python >= 3.10, NumPy < 2.0, Jetson Orin Nano
|
||||
- All existing tests (71) must pass after each task
|
||||
- Follow existing Flat Design system (teal/orange)
|
||||
- Git remote: `git.proit.id/andrew/feedmill-recounter.git`
|
||||
|
||||
## File Structure
|
||||
|
||||
| Action | File | Responsibility |
|
||||
|--------|------|----------------|
|
||||
| Modify | `src/pipeline.py` | Class filter before counter, switch to MultiClassLineCounter, box counts in results |
|
||||
| Modify | `src/job.py` | Box count fields in JobResult, pass from pipeline results |
|
||||
| Modify | `src/dashboard.py` | Label from det.class_name not hardcoded "sack" |
|
||||
| Modify | `templates/status.html` | Show sack vs box counts separately |
|
||||
| Modify | `app.py` | Download route compresses via ffmpeg to 250MB |
|
||||
| Modify | `tests/` | New tests for filter, box counts, compression |
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Fix truck counting bug (pipeline + dashboard)
|
||||
|
||||
**Files:**
|
||||
- Modify: `src/pipeline.py`
|
||||
- Modify: `src/dashboard.py`
|
||||
- Create/Modify: `tests/test_pipeline.py` (add filter tests)
|
||||
|
||||
**Root cause:**
|
||||
1. `run_merged_pipeline` line 519: `counter.update([d for d in tracked if d.track_id is not None])` — no class filter, trucks counted
|
||||
2. `dashboard.py` line 155: `label = "sack"` hardcoded
|
||||
3. Both pipelines use `LineCrossCounter` which counts any class
|
||||
|
||||
**Changes:**
|
||||
|
||||
1. `src/pipeline.py` — import `MultiClassLineCounter` instead of / in addition to `LineCrossCounter`
|
||||
|
||||
2. `run_pipeline` (line 163): replace `LineCrossCounter(...)` with `MultiClassLineCounter(...)` (same args)
|
||||
|
||||
3. `run_merged_pipeline` (line 426): same replacement
|
||||
|
||||
4. `run_pipeline` line 248: `counter.update(tracked_sacks)` → filter first:
|
||||
```python
|
||||
counter.update([d for d in tracked_sacks if d.class_name in ("sack", "box")])
|
||||
```
|
||||
|
||||
5. `run_merged_pipeline` line 519:
|
||||
```python
|
||||
countable = [d for d in tracked if d.track_id is not None and d.class_name in ("sack", "box")]
|
||||
counter.update(countable)
|
||||
```
|
||||
|
||||
6. `src/dashboard.py` `_draw_detections` line 155: `label = "sack"` → `label = det.class_name`
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: `MultiClassLineCounter` from src/counting.py (exists, has `box_loading_count`, `box_unloading_count` properties)
|
||||
- Produces: counter objects are now `MultiClassLineCounter` — same `.loading_count`/`.unloading_count` properties plus new `.box_loading_count`/`.box_unloading_count`
|
||||
|
||||
**Steps:**
|
||||
- [ ] Write failing test: truck class detections not counted
|
||||
- [ ] Run test to verify fail
|
||||
- [ ] Apply 6 changes above
|
||||
- [ ] Run full suite: `python -m pytest tests/ --tb=short`
|
||||
- [ ] Commit: `fix: exclude truck class from counting, use MultiClassLineCounter, real class labels`
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Box count propagation to results UI
|
||||
|
||||
**Files:**
|
||||
- Modify: `src/pipeline.py` (PipelineResult, MergedPipelineResult, return statements)
|
||||
- Modify: `src/job.py` (JobResult, both branches of _run_job)
|
||||
- Modify: `templates/status.html`
|
||||
- Modify: `static/app.js` (live stats if applicable)
|
||||
|
||||
**Changes:**
|
||||
|
||||
1. `PipelineResult` dataclass: add `box_loading_count: int = 0`, `box_unloading_count: int = 0`, property `box_net_count`
|
||||
|
||||
2. `MergedPipelineResult` dataclass: same fields
|
||||
|
||||
3. `run_pipeline` return (line ~323): add `box_loading_count=counter.box_loading_count, box_unloading_count=counter.box_unloading_count`
|
||||
|
||||
4. `run_merged_pipeline` return (line ~585): same
|
||||
|
||||
5. `JobResult` dataclass: add `box_loading_count: int = 0`, `box_unloading_count: int = 0`, `box_net_count: int = 0`
|
||||
|
||||
6. Both `JobResult` construction sites in `_run_job` (lines ~220 and ~264): pass box counts from result
|
||||
|
||||
7. `templates/status.html` results card: add box stats blocks:
|
||||
```html
|
||||
<div class="stat">
|
||||
<span class="stat-value">{{ r.box_loading_count }}</span>
|
||||
<span class="stat-label">Box In</span>
|
||||
</div>
|
||||
<div class="stat">
|
||||
<span class="stat-value">{{ r.box_unloading_count }}</span>
|
||||
<span class="stat-label">Box Out</span>
|
||||
</div>
|
||||
```
|
||||
Keep existing Loading/Unloading as sack counts (rename labels to "Sack In"/"Sack Out" for clarity).
|
||||
|
||||
8. `/api/jobs/<job_id>` in app.py: include box counts in results JSON
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: box counts from Task 1's MultiClassLineCounter
|
||||
- Produces: `box_loading_count` etc. on JobResult — status.html and API consumers
|
||||
|
||||
**Steps:**
|
||||
- [ ] Write failing test for box count propagation
|
||||
- [ ] Run to verify fail
|
||||
- [ ] Apply changes
|
||||
- [ ] Run full suite
|
||||
- [ ] Commit: `feat: propagate sack vs box counts to results and UI`
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Download compression to max 250MB
|
||||
|
||||
**Files:**
|
||||
- Modify: `app.py` (download route)
|
||||
- Create: `tests/test_download_compress.py`
|
||||
|
||||
**Changes:**
|
||||
|
||||
1. Add helper `_compress_for_download(file_path, max_bytes=250*1024*1024) -> str | None`:
|
||||
- If file size <= max_bytes, return None (no compression)
|
||||
- Probe duration via ffprobe
|
||||
- Calculate target bitrate: `int((max_bytes * 8 * 0.92) / duration)` (8% safety margin, avoid container overhead overshoot)
|
||||
- ffmpeg: `ffmpeg -y -i input -c:v libx264 -b:v {bitrate} -c:a aac -b:a 64k output.mp4`
|
||||
- Output to `<job_dir>/<name>_compressed.mp4`
|
||||
- If compressed file exists, reuse it
|
||||
- Return compressed path or None
|
||||
|
||||
2. `download()` route: call helper, send compressed file if produced, else original
|
||||
|
||||
3. Handle ffmpeg missing: fall back to original file (no error)
|
||||
|
||||
**Steps:**
|
||||
- [ ] Write failing tests (small file not compressed; helper signature)
|
||||
- [ ] Run to verify fail
|
||||
- [ ] Implement helper + route change
|
||||
- [ ] Run full suite
|
||||
- [ ] Commit: `feat: compress downloads to max 250MB via ffmpeg`
|
||||
+1
-1
@@ -152,7 +152,7 @@ class DashboardOverlay:
|
||||
) -> None:
|
||||
for det in detections:
|
||||
x1, y1, x2, y2 = [int(v) for v in det.bbox]
|
||||
label = "sack"
|
||||
label = det.class_name
|
||||
if det.track_id is not None:
|
||||
label += f" #{det.track_id}"
|
||||
label += f" {det.confidence:.0%}"
|
||||
|
||||
+10
-5
@@ -13,7 +13,7 @@ import cv2
|
||||
from ultralytics import YOLO
|
||||
|
||||
from src.batch import BatchLifecycleManager
|
||||
from src.counting import LineCrossCounter
|
||||
from src.counting import MultiClassLineCounter
|
||||
from src.dashboard import DashboardOverlay
|
||||
from src.detection import BaseDetector
|
||||
from src.interfaces import Detection
|
||||
@@ -72,6 +72,11 @@ def apply_class_filter(
|
||||
return [d for d in detections if d.class_name in class_filter]
|
||||
|
||||
|
||||
def countable_detections(detections: list[Detection]) -> list[Detection]:
|
||||
"""Keep only classes the counter accepts (sack/box); drop truck etc."""
|
||||
return [d for d in detections if d.class_name in ("sack", "box")]
|
||||
|
||||
|
||||
def run_pipeline(
|
||||
video_path: str,
|
||||
model_config: ModelConfig,
|
||||
@@ -160,7 +165,7 @@ def run_pipeline(
|
||||
tracker = ByteTrackTracker(shared_model, conf=sack_conf)
|
||||
stabilizer = BboxStabilizer()
|
||||
roi_tracker = TruckROITracker(frame_width=w, frame_height=h)
|
||||
counter = LineCrossCounter(
|
||||
counter = MultiClassLineCounter(
|
||||
line_y=int(h * 0.50),
|
||||
line_x_start=int(w * 0.38),
|
||||
line_x_end=int(w * 0.72),
|
||||
@@ -245,7 +250,7 @@ def run_pipeline(
|
||||
else:
|
||||
tracked_sacks = stable
|
||||
|
||||
counter.update(tracked_sacks)
|
||||
counter.update(countable_detections(tracked_sacks))
|
||||
|
||||
# Annotate frame
|
||||
viz = dashboard.draw(
|
||||
@@ -423,7 +428,7 @@ def run_merged_pipeline(
|
||||
tracker = ByteTrackTracker(tracker_model, conf=sack_conf)
|
||||
stabilizer = BboxStabilizer()
|
||||
roi_tracker = TruckROITracker(frame_width=w, frame_height=h)
|
||||
counter = LineCrossCounter(
|
||||
counter = MultiClassLineCounter(
|
||||
line_y=int(h * 0.50),
|
||||
line_x_start=int(w * 0.38),
|
||||
line_x_end=int(w * 0.72),
|
||||
@@ -516,7 +521,7 @@ def run_merged_pipeline(
|
||||
tracked = stable
|
||||
|
||||
# Count using only detections with track_id (tracker results)
|
||||
counter.update([d for d in tracked if d.track_id is not None])
|
||||
counter.update(countable_detections([d for d in tracked if d.track_id is not None]))
|
||||
else:
|
||||
tracked = []
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ import os
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pytest
|
||||
from src.counting import MultiClassLineCounter
|
||||
from src.pipeline import apply_class_filter, run_pipeline, PipelineResult
|
||||
from src.interfaces import Detection
|
||||
from src.model_registry import ModelConfig
|
||||
@@ -109,3 +110,81 @@ def test_cancel_check_breaks_pipeline(tmp_path):
|
||||
import inspect
|
||||
sig = inspect.signature(run_pipeline)
|
||||
assert "cancel_check" in sig.parameters
|
||||
|
||||
|
||||
# ── Truck class must never be counted ────────────────────────────────────
|
||||
|
||||
|
||||
def _cross_det(class_name: str, track_id: int, y1: float, x: float) -> Detection:
|
||||
"""Detection whose top edge (y1) sits above/below a line at y=100."""
|
||||
return Detection(
|
||||
bbox=(x, y1, x + 40, y1 + 40),
|
||||
confidence=0.9,
|
||||
class_id=0,
|
||||
class_name=class_name,
|
||||
track_id=track_id,
|
||||
)
|
||||
|
||||
|
||||
def test_countable_detections_excludes_truck():
|
||||
"""Pipeline's count filter keeps sack/box, drops truck."""
|
||||
from src.pipeline import countable_detections
|
||||
|
||||
dets = [
|
||||
_cross_det("sack", 1, 50.0, 100.0),
|
||||
_cross_det("truck", 2, 50.0, 200.0),
|
||||
_cross_det("box", 3, 50.0, 150.0),
|
||||
]
|
||||
assert [d.class_name for d in countable_detections(dets)] == ["sack", "box"]
|
||||
|
||||
|
||||
def test_truck_crossing_not_counted_after_filter():
|
||||
"""Sack + truck both cross the line; only the sack increments the counter."""
|
||||
from src.pipeline import countable_detections
|
||||
|
||||
counter = MultiClassLineCounter(line_y=100, line_x_start=0, line_x_end=320, margin=20)
|
||||
# frame 1: both above the line
|
||||
counter.update(countable_detections([
|
||||
_cross_det("sack", 1, 50.0, 100.0),
|
||||
_cross_det("truck", 2, 50.0, 200.0),
|
||||
]))
|
||||
# frame 2: both below the line (crossing)
|
||||
events = counter.update(countable_detections([
|
||||
_cross_det("sack", 1, 150.0, 100.0),
|
||||
_cross_det("truck", 2, 150.0, 200.0),
|
||||
]))
|
||||
|
||||
assert counter.loading_count == 1
|
||||
assert counter.unloading_count == 0
|
||||
assert [ev["track_id"] for ev in events] == [1]
|
||||
assert all(ev["track_id"] != 2 for ev in events)
|
||||
|
||||
|
||||
def test_multiclass_line_counter_ignores_truck_class():
|
||||
"""MultiClassLineCounter drops truck detections entirely (defense in depth)."""
|
||||
counter = MultiClassLineCounter(line_y=100, line_x_start=0, line_x_end=320, margin=20)
|
||||
counter.update([_cross_det("truck", 9, 50.0, 100.0)])
|
||||
events = counter.update([_cross_det("truck", 9, 150.0, 100.0)])
|
||||
|
||||
assert events == []
|
||||
assert counter.loading_count == 0
|
||||
assert counter.unloading_count == 0
|
||||
assert counter.box_loading_count == 0
|
||||
assert counter.box_unloading_count == 0
|
||||
assert counter.net_count == 0
|
||||
|
||||
|
||||
def test_dashboard_label_uses_class_name(monkeypatch):
|
||||
"""_draw_detections renders det.class_name, not hardcoded 'sack'."""
|
||||
from src import dashboard as dash_mod
|
||||
|
||||
drawn: list[str] = []
|
||||
monkeypatch.setattr(
|
||||
dash_mod.cv2, "putText",
|
||||
lambda frame, text, *args, **kwargs: drawn.append(text),
|
||||
)
|
||||
frame = np.zeros((240, 320, 3), dtype=np.uint8)
|
||||
dash_mod.DashboardOverlay()._draw_detections(
|
||||
frame, [_cross_det("truck", 7, 50.0, 100.0)]
|
||||
)
|
||||
assert any(t.startswith("truck") for t in drawn), drawn
|
||||
Reference in new issue
Block a user