feat(PPO_Bot): command abstraction layer — goto/aimTo controllers (#24)
- Add gotoTick/aimToTick controller functions (#25) - Update network dims: actor 5→6, state 42→44 (#26) - Rewrite mapActions for 6-dim command space (#27) - Delete stale weight files (shape mismatch) - Fix existing tests for new signatures Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+35
-23
@@ -2,32 +2,44 @@
|
||||
|
||||
import arraymancer
|
||||
import std/math
|
||||
import ./controllers
|
||||
|
||||
func sigmoid(x: float): float = 1.0 / (1.0 + exp(-x))
|
||||
|
||||
type
|
||||
BotActions* = object
|
||||
targetSpeed*: float32
|
||||
turnRate*: float32
|
||||
gunTurnRate*: float32
|
||||
targetSpeed*: float
|
||||
turnRate*: float
|
||||
gunTurnRate*: float
|
||||
shouldFire*: bool
|
||||
firePower*: float32
|
||||
firePower*: float
|
||||
gotoX*: float
|
||||
gotoY*: float
|
||||
aimToX*: float
|
||||
aimToY*: float
|
||||
|
||||
proc mapActions*(rawActions: Tensor[float32], currentSpeed: float32, gunHeat: float32): BotActions =
|
||||
## rawActions: [5] tensor from actorForward.
|
||||
## Ranges:
|
||||
## targetSpeed ∈ [-8, 8]
|
||||
## turnRate ∈ [-(10 - 0.75|v|), (10 - 0.75|v|)]
|
||||
## gunTurnRate ∈ [-20, 20]
|
||||
## firePower ∈ [0.1, 3.0]
|
||||
let r0 = tanh(rawActions[0].float64).float32
|
||||
let r1 = tanh(rawActions[1].float64).float32
|
||||
let r2 = tanh(rawActions[2].float64).float32
|
||||
let r3 = tanh(rawActions[3].float64).float32
|
||||
let r4 = rawActions[4]
|
||||
proc mapActions*(rawActions: Tensor[float32],
|
||||
gunHeat: float,
|
||||
arenaWidth, arenaHeight: float,
|
||||
botX, botY, direction, speed, gunDirection: float): BotActions =
|
||||
## rawActions: [6] tensor from actorForward.
|
||||
## Dims 0–1: goto x/y, 2–3: aimTo x/y, 4: fire decision, 5: fire power.
|
||||
let gotoX = sigmoid(rawActions[0].float) * arenaWidth
|
||||
let gotoY = sigmoid(rawActions[1].float) * arenaHeight
|
||||
let aimToX = sigmoid(rawActions[2].float) * arenaWidth
|
||||
let aimToY = sigmoid(rawActions[3].float) * arenaHeight
|
||||
let fireDec = tanh(rawActions[4].float)
|
||||
let fp = sigmoid(rawActions[5].float) * 2.9 + 0.1
|
||||
|
||||
result.targetSpeed = r0 * 8.0'f32
|
||||
result.turnRate = r1 * (10.0'f32 - 0.75'f32 * abs(currentSpeed))
|
||||
result.gunTurnRate = r2 * 20.0'f32
|
||||
let fireDecision = r3
|
||||
result.shouldFire = fireDecision >= 0.0'f32 and gunHeat <= 0.0'f32
|
||||
# sigmoid(r4) * 2.9 + 0.1 → [0.1, 3.0]
|
||||
result.firePower = (1.0'f32 / (1.0'f32 + exp(-r4.float64).float32)) * 2.9'f32 + 0.1'f32
|
||||
let (ts, tr) = gotoTick(gotoX, gotoY, botX, botY, direction, speed)
|
||||
let gtr = aimToTick(aimToX, aimToY, botX, botY, gunDirection)
|
||||
|
||||
result.gotoX = gotoX
|
||||
result.gotoY = gotoY
|
||||
result.aimToX = aimToX
|
||||
result.aimToY = aimToY
|
||||
result.targetSpeed = ts
|
||||
result.turnRate = tr
|
||||
result.gunTurnRate = gtr
|
||||
result.shouldFire = fireDec >= 0.0 and gunHeat <= 0.0
|
||||
result.firePower = fp
|
||||
|
||||
Reference in New Issue
Block a user