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