diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/actions.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/actions.nim new file mode 100644 index 0000000..d4b10c3 --- /dev/null +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/actions.nim @@ -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 diff --git a/SAC_LSTM_Bot/tests/test_actions.nim b/SAC_LSTM_Bot/tests/test_actions.nim new file mode 100644 index 0000000..498d8b4 --- /dev/null +++ b/SAC_LSTM_Bot/tests/test_actions.nim @@ -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