amanmurari commited on
Commit
eef9e42
·
verified ·
1 Parent(s): de9d04c

Upload folder using huggingface_hub

Browse files
Files changed (1) hide show
  1. inference.py +11 -3
inference.py CHANGED
@@ -76,8 +76,16 @@ def _heuristic_phase(obs: TrafficObservation, task: str) -> int:
76
  ns_urg = max(obs.emergency_urgency[0], obs.emergency_urgency[1])
77
  ew_urg = max(obs.emergency_urgency[2], obs.emergency_urgency[3])
78
  cur = obs.current_phase
 
 
 
 
 
 
 
 
79
 
80
- # Critical emergency
81
  if ns_em > 0 and ns_urg >= 8 and ew_em > 0 and ew_urg >= 8:
82
  return 2
83
  if ns_em > 0 and ns_urg >= 8:
@@ -85,7 +93,7 @@ def _heuristic_phase(obs: TrafficObservation, task: str) -> int:
85
  if ew_em > 0 and ew_urg >= 8:
86
  return 1
87
 
88
- # Moderate emergency
89
  if ns_em > 0 and ns_urg >= 5:
90
  if ew_em == 0 or ns_urg >= ew_urg:
91
  return 0
@@ -98,7 +106,7 @@ def _heuristic_phase(obs: TrafficObservation, task: str) -> int:
98
  if cur in (0, 3): return 0
99
  if cur in (1, 4): return 1
100
 
101
- # Pressure
102
  ns_p, ew_p = _compute_pressures(obs)
103
  ratio = 1.5 if task == "basic_flow" else 1.2
104
  if ns_p > ew_p * ratio: return 0
 
76
  ns_urg = max(obs.emergency_urgency[0], obs.emergency_urgency[1])
77
  ew_urg = max(obs.emergency_urgency[2], obs.emergency_urgency[3])
78
  cur = obs.current_phase
79
+ ns_q = obs.queue_lengths[0] + obs.queue_lengths[1]
80
+ ew_q = obs.queue_lengths[2] + obs.queue_lengths[3]
81
+
82
+ # Collision risk — rotate to larger queue to drain before gridlock (-200 penalty)
83
+ total_q = sum(obs.queue_lengths)
84
+ if total_q > 28 and obs.time_in_phase > 14:
85
+ if cur in (0, 3): return 1 if ew_q > ns_q else 0
86
+ return 0 if ns_q > ew_q else 1
87
 
88
+ # Critical emergency (urgency >= 8)
89
  if ns_em > 0 and ns_urg >= 8 and ew_em > 0 and ew_urg >= 8:
90
  return 2
91
  if ns_em > 0 and ns_urg >= 8:
 
93
  if ew_em > 0 and ew_urg >= 8:
94
  return 1
95
 
96
+ # Moderate emergency (urgency >= 5)
97
  if ns_em > 0 and ns_urg >= 5:
98
  if ew_em == 0 or ns_urg >= ew_urg:
99
  return 0
 
106
  if cur in (0, 3): return 0
107
  if cur in (1, 4): return 1
108
 
109
+ # Pressure ratio
110
  ns_p, ew_p = _compute_pressures(obs)
111
  ratio = 1.5 if task == "basic_flow" else 1.2
112
  if ns_p > ew_p * ratio: return 0