95 lines
3.3 KiB
Go
95 lines
3.3 KiB
Go
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
|
|
}
|