feat(SAC_LSTM_Bot): action mapping module (#43)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,36 @@
|
|||||||
|
## actions.nim — map raw network output (4 tanh values) to bot intent fields.
|
||||||
|
##
|
||||||
|
## Note on acceleration vs targetSpeed:
|
||||||
|
## TankRoyale uses setTargetSpeed(), not setAcceleration().
|
||||||
|
## The mapped `acceleration` field is a delta; callers must compute:
|
||||||
|
## newTargetSpeed = clamp(currentSpeed + acceleration, -8.0, 8.0)
|
||||||
|
## and call setTargetSpeed(newTargetSpeed).
|
||||||
|
|
||||||
|
import arraymancer
|
||||||
|
|
||||||
|
const ACTION_DIM* = 4
|
||||||
|
|
||||||
|
type
|
||||||
|
MappedActions* = object
|
||||||
|
turnRate*: float ## degrees/tick, speed-aware; [-10, 10] at speed 0
|
||||||
|
acceleration*: float ## delta speed in [-2, +1]; caller adds to currentSpeed
|
||||||
|
gunTurnRate*: float ## degrees/tick in [-20, 20]
|
||||||
|
firePower*: float ## 0 = don't fire; (0.1, 3.0] = fire with this power
|
||||||
|
|
||||||
|
proc mapActions*(networkOutput: Tensor[float32],
|
||||||
|
currentSpeed: float,
|
||||||
|
gunHeat: float): MappedActions =
|
||||||
|
## networkOutput: [4] tensor of tanh values in [-1, 1].
|
||||||
|
let a0 = networkOutput[0].float
|
||||||
|
let a1 = networkOutput[1].float
|
||||||
|
let a2 = networkOutput[2].float
|
||||||
|
let a3 = networkOutput[3].float
|
||||||
|
|
||||||
|
result.turnRate = a0 * (10.0 - 0.75 * abs(currentSpeed))
|
||||||
|
# asymmetric accel: [-1,1] -> [-2, +1] via (value * 1.5 - 0.5)
|
||||||
|
result.acceleration = a1 * 1.5 - 0.5
|
||||||
|
result.gunTurnRate = a2 * 20.0
|
||||||
|
if a3 > 0.0 and gunHeat <= 0.0:
|
||||||
|
result.firePower = a3 * 2.9 + 0.1
|
||||||
|
else:
|
||||||
|
result.firePower = 0.0
|
||||||
@@ -0,0 +1,97 @@
|
|||||||
|
import unittest
|
||||||
|
import arraymancer
|
||||||
|
import std/math
|
||||||
|
import SAC_LSTM_Bot/actions
|
||||||
|
|
||||||
|
proc makeOutput(a0, a1, a2, a3: float): Tensor[float32] =
|
||||||
|
result = newTensor[float32](4)
|
||||||
|
result[0] = a0.float32
|
||||||
|
result[1] = a1.float32
|
||||||
|
result[2] = a2.float32
|
||||||
|
result[3] = a3.float32
|
||||||
|
|
||||||
|
suite "mapActions":
|
||||||
|
|
||||||
|
test "ACTION_DIM is 4":
|
||||||
|
check ACTION_DIM == 4
|
||||||
|
|
||||||
|
# Speed-aware turn rate
|
||||||
|
test "turn: output +1 at speed 0 -> +10 degrees":
|
||||||
|
let m = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.turnRate - 10.0) < 1e-6
|
||||||
|
|
||||||
|
test "turn: output -1 at speed 0 -> -10 degrees":
|
||||||
|
let m = mapActions(makeOutput(-1.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.turnRate - (-10.0)) < 1e-6
|
||||||
|
|
||||||
|
test "turn: output +1 at speed 8 -> +4 degrees":
|
||||||
|
let m = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||||
|
check abs(m.turnRate - 4.0) < 1e-6
|
||||||
|
|
||||||
|
test "turn: output -1 at speed 8 -> -4 degrees":
|
||||||
|
let m = mapActions(makeOutput(-1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||||
|
check abs(m.turnRate - (-4.0)) < 1e-6
|
||||||
|
|
||||||
|
test "turn: output 0 -> 0 regardless of speed":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 5.0, 0.0)
|
||||||
|
check abs(m.turnRate) < 1e-6
|
||||||
|
|
||||||
|
# Acceleration range
|
||||||
|
test "accel: output -1 -> -2.0":
|
||||||
|
let m = mapActions(makeOutput(0.0, -1.0, 0.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.acceleration - (-2.0)) < 1e-6
|
||||||
|
|
||||||
|
test "accel: output +1 -> +1.0":
|
||||||
|
let m = mapActions(makeOutput(0.0, 1.0, 0.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.acceleration - 1.0) < 1e-6
|
||||||
|
|
||||||
|
test "accel: output 0 -> midpoint -0.5":
|
||||||
|
# value*1.5 - 0.5 at value=0 -> -0.5 (correct midpoint between -2 and +1)
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.acceleration - (-0.5)) < 1e-6
|
||||||
|
|
||||||
|
# Gun turn rate
|
||||||
|
test "gun turn: output +1 -> +20 degrees":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 1.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.gunTurnRate - 20.0) < 1e-6
|
||||||
|
|
||||||
|
test "gun turn: output -1 -> -20 degrees":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, -1.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.gunTurnRate - (-20.0)) < 1e-6
|
||||||
|
|
||||||
|
test "gun turn: output 0 -> 0":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||||
|
check abs(m.gunTurnRate) < 1e-6
|
||||||
|
|
||||||
|
# Fire threshold
|
||||||
|
test "fire: output -0.5 -> no fire (firePower == 0)":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -0.5), 0.0, 0.0)
|
||||||
|
check m.firePower == 0.0
|
||||||
|
|
||||||
|
test "fire: output 0.0 -> no fire (boundary, not positive)":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.0), 0.0, 0.0)
|
||||||
|
check m.firePower == 0.0
|
||||||
|
|
||||||
|
test "fire: output +0.5 -> fire with correct power":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.5), 0.0, 0.0)
|
||||||
|
# 0.5 * 2.9 + 0.1 = 1.55
|
||||||
|
check abs(m.firePower - 1.55) < 1e-5
|
||||||
|
|
||||||
|
test "fire: output +1.0 -> fire power near 3.0":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 1.0), 0.0, 0.0)
|
||||||
|
check abs(m.firePower - 3.0) < 1e-5
|
||||||
|
|
||||||
|
test "fire: output +1.0 but gunHeat > 0 -> no fire":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 1.0), 0.0, 1.5)
|
||||||
|
check m.firePower == 0.0
|
||||||
|
|
||||||
|
test "fire: output +0.001 (just above 0) -> fires with power near 0.1":
|
||||||
|
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.001), 0.0, 0.0)
|
||||||
|
check m.firePower > 0.0
|
||||||
|
check m.firePower < 0.2
|
||||||
|
|
||||||
|
# Speed-aware turn with negative speed (reverse)
|
||||||
|
test "turn: speed -8 (reversing) -> same magnitude as speed +8":
|
||||||
|
let fwd = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||||
|
let rev = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), -8.0, 0.0)
|
||||||
|
check abs(fwd.turnRate - rev.turnRate) < 1e-6
|
||||||
Reference in New Issue
Block a user