feat(v2): add fall event regression engine
This commit is contained in:
@@ -0,0 +1,94 @@
|
||||
package fall
|
||||
|
||||
import "fmt"
|
||||
|
||||
type Engine struct {
|
||||
config EngineConfig
|
||||
tracker *tracker
|
||||
policy *policy
|
||||
stateMachine *stateMachine
|
||||
previousEvidence map[string]PoseEvidence
|
||||
activeTrackIDs map[string]bool
|
||||
}
|
||||
|
||||
func NewEngine(config EngineConfig) (*Engine, error) {
|
||||
if config.KeypointConfidenceThreshold < 0 || config.KeypointConfidenceThreshold > 1 {
|
||||
return nil, fmt.Errorf("keypoint confidence threshold must be between 0 and 1")
|
||||
}
|
||||
if config.SuspectWindowSeconds < 0 {
|
||||
return nil, fmt.Errorf("suspect window seconds must be non-negative")
|
||||
}
|
||||
if config.HorizontalAngleThresholdDegrees < 0 || config.HorizontalAngleThresholdDegrees > 90 {
|
||||
return nil, fmt.Errorf("horizontal angle threshold degrees must be between 0 and 90")
|
||||
}
|
||||
tracker, err := newTracker(0.2, 2.0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stateMachine, err := newStateMachine(config)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Engine{
|
||||
config: config, tracker: tracker, policy: newPolicy(config.SuspectWindowSeconds, config.RequireRapidDrop),
|
||||
stateMachine: stateMachine, previousEvidence: make(map[string]PoseEvidence), activeTrackIDs: make(map[string]bool),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (engine *Engine) StateOf(trackID string) State {
|
||||
return engine.stateMachine.stateOf(trackID)
|
||||
}
|
||||
|
||||
func (engine *Engine) Process(frame Frame) FrameResult {
|
||||
tracked, err := engine.tracker.update(frame.Poses, frame.Timestamp, frame.Width, frame.Height)
|
||||
if err != nil {
|
||||
return FrameResult{}
|
||||
}
|
||||
currentIDs := make(map[string]bool, len(tracked))
|
||||
for _, trackedPose := range tracked {
|
||||
currentIDs[trackedPose.TrackID] = true
|
||||
}
|
||||
events := engine.rejectMissingTracks(currentIDs, frame.Timestamp)
|
||||
people := make([]PersonAnalysis, 0, len(tracked))
|
||||
for _, trackedPose := range tracked {
|
||||
quality := assessPoseQuality(trackedPose.Pose, engine.config.KeypointConfidenceThreshold, engine.config.RequireLowerBody)
|
||||
var previous *PoseEvidence
|
||||
if candidate, found := engine.previousEvidence[trackedPose.TrackID]; found {
|
||||
previous = &candidate
|
||||
}
|
||||
poseEvidence := extractEvidence(trackedPose.Pose, quality, previous, engine.config.HorizontalAngleThresholdDegrees)
|
||||
stateBefore := engine.stateMachine.stateOf(trackedPose.TrackID)
|
||||
evidence := engine.policy.evaluate(trackedPose.TrackID, poseEvidence, frame.Timestamp, stateBefore)
|
||||
newEvents, err := engine.stateMachine.update(trackedPose.TrackID, evidence, frame.Timestamp)
|
||||
if err == nil {
|
||||
events = append(events, newEvents...)
|
||||
}
|
||||
if poseEvidence.Accepted {
|
||||
engine.previousEvidence[trackedPose.TrackID] = poseEvidence
|
||||
} else {
|
||||
delete(engine.previousEvidence, trackedPose.TrackID)
|
||||
}
|
||||
people = append(people, PersonAnalysis{
|
||||
TrackedPose: trackedPose, PoseEvidence: poseEvidence, Evidence: evidence,
|
||||
State: engine.stateMachine.stateOf(trackedPose.TrackID),
|
||||
})
|
||||
}
|
||||
engine.activeTrackIDs = currentIDs
|
||||
return FrameResult{People: people, Events: events}
|
||||
}
|
||||
|
||||
func (engine *Engine) rejectMissingTracks(currentIDs map[string]bool, now float64) []Event {
|
||||
events := make([]Event, 0)
|
||||
for trackID := range engine.activeTrackIDs {
|
||||
if currentIDs[trackID] {
|
||||
continue
|
||||
}
|
||||
delete(engine.previousEvidence, trackID)
|
||||
engine.policy.evaluate(trackID, PoseEvidence{}, now, engine.stateMachine.stateOf(trackID))
|
||||
newEvents, err := engine.stateMachine.update(trackID, Evidence{}, now)
|
||||
if err == nil {
|
||||
events = append(events, newEvents...)
|
||||
}
|
||||
}
|
||||
return events
|
||||
}
|
||||
Reference in New Issue
Block a user