j130 learned movement: outcome label (P(hit|state,g)) mode + Gate A; pre-registered outcome arms
This commit is contained in:
@@ -62,10 +62,19 @@
|
||||
## TR_LEARNED_WALL_MARGIN wall margin (48 px)
|
||||
## TR_LEARNED_GLOBAL =1: ignore the state (ABLATION, the same mover
|
||||
## with a pure global histogram)
|
||||
## TR_LEARNED_LABEL `histogram` (default) or `outcome` (job j130).
|
||||
## `outcome` replaces the label (the GF bin we
|
||||
## cross) with the OUTCOME: a counted SBC over
|
||||
## (state, candidate bin) estimating P(hit | state,
|
||||
## g), trained on the dense label `hit AND |g-b| <=
|
||||
## window` (b = the bin the wave resolved at, window
|
||||
## = the bot radius as an angle at the wave's
|
||||
## distance). See docs/movement_campaign.md,
|
||||
## "outcome label".
|
||||
## TR_LEARNED_LOG per-decision log line
|
||||
|
||||
import std/[math, os]
|
||||
from std/strutils import parseFloat, parseInt, strip
|
||||
from std/strutils import parseFloat, parseInt, strip, toLowerAscii
|
||||
import gun_harness/gun_interface
|
||||
import movement_harness/movement_interface
|
||||
import bitbrain/sbc
|
||||
@@ -98,8 +107,19 @@ const
|
||||
LearnedRadialFracEnv* = "TR_LEARNED_RADIAL_FRAC"
|
||||
LearnedWallMarginEnv* = "TR_LEARNED_WALL_MARGIN"
|
||||
LearnedGlobalEnv* = "TR_LEARNED_GLOBAL"
|
||||
LearnedLabelEnv* = "TR_LEARNED_LABEL"
|
||||
LearnedLogEnv* = "TR_LEARNED_LOG"
|
||||
|
||||
## Number of joint (vlat,dist,room,turn) state codes = 4^4.
|
||||
LS_STATES = LS_Q * LS_Q * LS_Q * LS_Q
|
||||
## Prior-mix weight for the 2-class outcome readout.
|
||||
OutcomePriorAlpha = 1.0
|
||||
|
||||
type
|
||||
LearnedLabelMode* = enum
|
||||
llHistogram ## label = the GF bin the wave crossed us at (j128)
|
||||
llOutcome ## label = the dense hit outcome (job j130)
|
||||
|
||||
var
|
||||
LearnedDecayEvery* = 128
|
||||
LearnedDecayShift* = 1
|
||||
@@ -111,6 +131,7 @@ var
|
||||
LearnedRadialFrac* = 0.35
|
||||
LearnedWallMargin* = 48.0
|
||||
LearnedGlobal* = false
|
||||
LearnedLabel* = llHistogram
|
||||
LearnedLog* = false
|
||||
|
||||
proc getEnvFloat(name: string, default: float): float =
|
||||
@@ -144,6 +165,10 @@ proc loadLearnedEnv*() =
|
||||
LearnedWallMargin = max(0.0, getEnvFloat(LearnedWallMarginEnv, 48.0))
|
||||
LearnedGlobal = envOn(LearnedGlobalEnv)
|
||||
LearnedLog = envOn(LearnedLogEnv)
|
||||
LearnedLabel =
|
||||
case getEnv(LearnedLabelEnv, "").strip().toLowerAscii()
|
||||
of "outcome": llOutcome
|
||||
else: llHistogram
|
||||
|
||||
loadLearnedEnv()
|
||||
|
||||
@@ -201,13 +226,17 @@ type
|
||||
power: float64
|
||||
ticksLeft: int ## ticks until the nominal arrival
|
||||
fresh: bool ## created this tick: do not age it yet
|
||||
selfEnergyAtFire: float64 ## our energy at the fire tick (hit detection)
|
||||
stateRow: int ## (vlat, dist) code
|
||||
stateCol: int ## (room, turn) code
|
||||
|
||||
LearnedSurferModule* = object
|
||||
sbc*: Sbc
|
||||
outcome*: Sbc ## llOutcome: (state, candidate bin) -> hit/miss
|
||||
glob*: array[LS_BINS, int]
|
||||
glc: int ## learns since the last global-histogram decay
|
||||
hitGlobal: int ## llOutcome: hits seen (for the prior)
|
||||
missGlobal: int ## llOutcome: non-hits seen (for the prior)
|
||||
scores: seq[float]
|
||||
waves: seq[LSWave]
|
||||
prevEnergy: seq[tuple[id: int, energy: float]]
|
||||
@@ -234,6 +263,8 @@ proc resetRound*(m: var LearnedSurferModule) =
|
||||
m.decisions = 0
|
||||
for i in 0..<LS_BINS: m.glob[i] = 0
|
||||
m.glc = 0
|
||||
m.hitGlobal = 0
|
||||
m.missGlobal = 0
|
||||
|
||||
proc initLearnedSurfer*(): LearnedSurferModule =
|
||||
result.debugGraphics = false
|
||||
@@ -242,12 +273,21 @@ proc initLearnedSurfer*(): LearnedSurferModule =
|
||||
max(0, LearnedDecayShift))
|
||||
result.sbc.decayEvery = max(0, LearnedDecayEvery)
|
||||
result.sbc.decayShift = max(0, LearnedDecayShift)
|
||||
# The outcome model learns LS_BINS samples per resolved wave, so its decay
|
||||
# period is scaled to the same WAVE cadence as the histogram's (decayEvery
|
||||
# waves), keeping the two labels' forgetting comparable.
|
||||
block:
|
||||
let oe = if LearnedDecayEvery > 0: LearnedDecayEvery * LS_BINS else: 0
|
||||
result.outcome = initCountedSbc(LS_STATES, 2, oe, max(0, LearnedDecayShift))
|
||||
result.outcome.decayEvery = oe
|
||||
result.outcome.decayShift = max(0, LearnedDecayShift)
|
||||
result.scores = newSeq[float](LS_BINS)
|
||||
result.resetRound()
|
||||
|
||||
proc resetBattle*(m: var LearnedSurferModule) =
|
||||
## Hard wipe of the learned memory (round the SBC's counters and `clear`).
|
||||
m.sbc.clear()
|
||||
m.outcome.clear()
|
||||
m.resetRound()
|
||||
|
||||
proc clearGraphics*(m: var LearnedSurferModule) {.inline.} = discard
|
||||
@@ -260,18 +300,57 @@ proc prior(m: LearnedSurferModule, k: int): float =
|
||||
for g in m.glob: tot += g
|
||||
(float(m.glob[k]) + 1.0) / (float(tot) + float(LS_BINS))
|
||||
|
||||
proc learnWave(m: var LearnedSurferModule, row, col, bin: int) =
|
||||
## One supervised sample: (state) -> (GF bin at arrival), in the counted SBC.
|
||||
## The SBC applies its own global decay every `decayEvery` learns; the global
|
||||
## histogram below is aged on the same schedule so the two stay comparable.
|
||||
discard m.sbc.learn([row.int32], [col.int32], bin)
|
||||
if m.glob[bin] < 255: inc m.glob[bin]
|
||||
inc m.glc
|
||||
if LearnedDecayEvery > 0 and LearnedDecayShift > 0 and
|
||||
m.glc >= LearnedDecayEvery:
|
||||
for k in 0..<LS_BINS:
|
||||
m.glob[k] = m.glob[k] - (m.glob[k] shr LearnedDecayShift)
|
||||
m.glc = 0
|
||||
proc learnWave(m: var LearnedSurferModule, w: LSWave, bin: int, hit: bool) =
|
||||
## One resolved wave -> training samples.
|
||||
## llHistogram: one sample, (state) -> (GF bin at arrival).
|
||||
## llOutcome: LS_BINS samples, (state, candidate g) -> hit, where the label
|
||||
## is `hit and |g - bin| <= window` (the dense counterfactual:
|
||||
## would this wave have hit us at candidate g?). `window` is
|
||||
## the bot radius as an angle at the wave's distance, in GF
|
||||
## bins — the angular/physical tolerance of the body. The SBC
|
||||
## applies its own global decay; the hit/miss counters below
|
||||
## are aged on the same schedule.
|
||||
case LearnedLabel
|
||||
of llHistogram:
|
||||
discard m.sbc.learn([w.stateRow.int32], [w.stateCol.int32], bin)
|
||||
if m.glob[bin] < 255: inc m.glob[bin]
|
||||
inc m.glc
|
||||
if LearnedDecayEvery > 0 and LearnedDecayShift > 0 and
|
||||
m.glc >= LearnedDecayEvery:
|
||||
for k in 0..<LS_BINS:
|
||||
m.glob[k] = m.glob[k] - (m.glob[k] shr LearnedDecayShift)
|
||||
m.glc = 0
|
||||
of llOutcome:
|
||||
let maxA = mea(w.speed)
|
||||
let wbin =
|
||||
if maxA > 1e-9:
|
||||
(arcsin(min(BotRadius / max(w.startDist, 1.0), 1.0)) / maxA) *
|
||||
float(LS_BINS - 1) / 2.0
|
||||
else: 0.0
|
||||
let sc = if LearnedGlobal: 0 else: w.stateRow * LS_Q * LS_Q + w.stateCol
|
||||
for g in 0..<LS_BINS:
|
||||
let lab = if hit and abs(float(g - bin)) <= wbin: 1 else: 0
|
||||
discard m.outcome.learn([sc.int32], [g.int32], lab)
|
||||
inc m.glc
|
||||
if hit: inc m.hitGlobal else: inc m.missGlobal
|
||||
if LearnedDecayEvery > 0 and LearnedDecayShift > 0 and
|
||||
m.glc >= LearnedDecayEvery * LS_BINS:
|
||||
m.hitGlobal = m.hitGlobal - (m.hitGlobal shr LearnedDecayShift)
|
||||
m.missGlobal = m.missGlobal - (m.missGlobal shr LearnedDecayShift)
|
||||
m.glc = 0
|
||||
|
||||
proc predictHit*(m: LearnedSurferModule, row, col, g: int): float =
|
||||
## P(hit | state, candidate bin g) — the `outcome` danger (lower = safer),
|
||||
## from the 2-class counted SBC read out with the per-cell posterior and
|
||||
## blended with the global hit rate.
|
||||
let t = m.hitGlobal + m.missGlobal
|
||||
let prior = if t == 0: 0.5 else: float(m.hitGlobal) / float(t)
|
||||
let sc = if LearnedGlobal: 0 else: row * LS_Q * LS_Q + col
|
||||
let c0 = m.outcome.countAt(sc, g, 0)
|
||||
let c1 = m.outcome.countAt(sc, g, 1)
|
||||
let n = c0 + c1
|
||||
if n == 0: return prior
|
||||
(float(c1) + OutcomePriorAlpha * prior) / (float(n) + OutcomePriorAlpha)
|
||||
|
||||
proc predictState*(m: var LearnedSurferModule, row, col: int,
|
||||
p: var array[LS_BINS, float]) =
|
||||
@@ -341,6 +420,7 @@ proc detectFire(m: var LearnedSurferModule, id: int, ex, ey, eenergy: float,
|
||||
startDist: d, power: drop,
|
||||
ticksLeft: max(1, int(ceil(d / max(bspeed, 1e-9)))),
|
||||
fresh: true,
|
||||
selfEnergyAtFire: ws.selfEnergy,
|
||||
stateRow: row, stateCol: col,
|
||||
)
|
||||
|
||||
@@ -371,7 +451,11 @@ proc computeMove*(m: var LearnedSurferModule, ws: WorldState): MoveCommand =
|
||||
if maxA >= 1e-9:
|
||||
let off = wrapPi(arctan2(botY - w.originY, botX - w.originX) - w.bearing)
|
||||
let bin = gfToBin(clamp(off / maxA, -1.0, 1.0))
|
||||
m.learnWave(w.stateRow, w.stateCol, bin)
|
||||
# llOutcome label: did THIS wave hit us? Our own energy dropped since
|
||||
# the fire tick. (One wave is live at a time in 1v1; ramming also drops
|
||||
# energy, so this is a proxy, not an oracle.)
|
||||
let hit = ws.selfEnergy < w.selfEnergyAtFire - 0.01
|
||||
m.learnWave(w, bin, hit)
|
||||
if LearnedLog:
|
||||
echo "[learned] resolve bin=", bin, " state=", w.stateRow, "/",
|
||||
w.stateCol, " d=", w.startDist.int, " e=", w.originX.int, ",",
|
||||
@@ -381,10 +465,13 @@ proc computeMove*(m: var LearnedSurferModule, ws: WorldState): MoveCommand =
|
||||
inc i
|
||||
|
||||
# ── 3. danger of every candidate bin, summed over every live wave ────────
|
||||
# llHistogram: precompute the predicted arrival-bin distribution per wave.
|
||||
# llOutcome: danger(wave, g) = P(hit | wave.state, g), computed on demand.
|
||||
var ps: seq[array[LS_BINS, float]]
|
||||
ps.setLen(m.waves.len)
|
||||
for j in 0..<m.waves.len:
|
||||
m.predictState(m.waves[j].stateRow, m.waves[j].stateCol, ps[j])
|
||||
if LearnedLabel == llHistogram:
|
||||
ps.setLen(m.waves.len)
|
||||
for j in 0..<m.waves.len:
|
||||
m.predictState(m.waves[j].stateRow, m.waves[j].stateCol, ps[j])
|
||||
|
||||
# nearest wave = largest radius/distance ratio (the one about to arrive)
|
||||
var nearest = -1
|
||||
@@ -433,7 +520,9 @@ proc computeMove*(m: var LearnedSurferModule, ws: WorldState): MoveCommand =
|
||||
let futureY = botY + uy * dodgeDist
|
||||
bestDirs[j] = if gfJ >= curGF: 1.0 else: -1.0
|
||||
|
||||
var danger = ps[nearest][j]
|
||||
var danger =
|
||||
if LearnedLabel == llHistogram: ps[nearest][j]
|
||||
else: m.predictHit(w.stateRow, w.stateCol, j)
|
||||
# every OTHER live wave contributes its own predicted mass at the bin the
|
||||
# candidate direction would land in for THAT wave.
|
||||
for k in 0..<m.waves.len:
|
||||
@@ -445,7 +534,8 @@ proc computeMove*(m: var LearnedSurferModule, ws: WorldState): MoveCommand =
|
||||
let off2 = wrapPi(arctan2(futureY - w2.originY, futureX - w2.originX) -
|
||||
w2.bearing)
|
||||
let b2 = gfToBin(clamp(off2 / mA2, -1.0, 1.0))
|
||||
danger += ps[k][b2]
|
||||
if LearnedLabel == llHistogram: danger += ps[k][b2]
|
||||
else: danger += m.predictHit(w2.stateRow, w2.stateCol, b2)
|
||||
|
||||
let wallHit = futureX < LearnedWallMargin or
|
||||
futureX > ws.arenaWidth - LearnedWallMargin or
|
||||
@@ -464,12 +554,22 @@ proc computeMove*(m: var LearnedSurferModule, ws: WorldState): MoveCommand =
|
||||
m.dir = m.strafeDir
|
||||
inc m.decisions
|
||||
if LearnedLog:
|
||||
var pAvg = 0.0
|
||||
for k in 0..<LS_BINS: pAvg += ps[nearest][k] / float(LS_BINS)
|
||||
var pAvg, pSafe, pCur: float
|
||||
if LearnedLabel == llHistogram:
|
||||
pAvg = 0.0
|
||||
for k in 0..<LS_BINS: pAvg += ps[nearest][k] / float(LS_BINS)
|
||||
pSafe = ps[nearest][bestBin]
|
||||
pCur = ps[nearest][gfToBin(curGF)]
|
||||
else:
|
||||
pAvg = 0.0
|
||||
for k in 0..<LS_BINS:
|
||||
pAvg += m.predictHit(w.stateRow, w.stateCol, k) / float(LS_BINS)
|
||||
pSafe = m.predictHit(w.stateRow, w.stateCol, bestBin)
|
||||
pCur = m.predictHit(w.stateRow, w.stateCol, gfToBin(curGF))
|
||||
echo "[learned] wave d=", d.int, " curGF=", (curGF * 100.0).int,
|
||||
" bestGF=", (bestGF * 100.0).int, " dir=", m.strafeDir.int,
|
||||
" pSafe=", (ps[nearest][bestBin] * 1000.0).int,
|
||||
" pCur=", (ps[nearest][gfToBin(curGF)] * 1000.0).int,
|
||||
" pSafe=", (pSafe * 1000.0).int,
|
||||
" pCur=", (pCur * 1000.0).int,
|
||||
" pAvg=", (pAvg * 1000.0).int,
|
||||
" state=", m.waves[nearest].stateRow, "/", m.waves[nearest].stateCol
|
||||
# perpendicular in the wave's own frame: +90 deg from origin->bot increases
|
||||
|
||||
Reference in New Issue
Block a user