feat(BNNBot): Hebbian weight matrix with virtual bullet learning — 870→7 bit forward pass
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -1,14 +1,22 @@
|
|||||||
# BNNBot — binary encoding data collector.
|
# BNNBot — Hebbian weight matrix with virtual bullet learning.
|
||||||
# Radar lock on enemy, encodes a 10-tick sliding window (870 bits), prints it.
|
# 870-bit input → 7-bit aim angle output via forward pass.
|
||||||
# No aiming, no firing — pure scan visualization for BNN research.
|
# Learns from virtual bullets (no real firing) via three-factor Hebbian rule.
|
||||||
|
|
||||||
import std/[math, os, strutils]
|
import std/[math, os, strutils, random]
|
||||||
import robocode_tankroyale_botapi
|
import robocode_tankroyale_botapi
|
||||||
import radar_lock/radar_lock as radar_lock
|
import radar_lock/radar_lock as radar_lock
|
||||||
import binary_encoding
|
import binary_encoding
|
||||||
|
import hebbian
|
||||||
|
|
||||||
const botJsonPath = currentSourcePath().parentDir / "BNNBot.json"
|
const botJsonPath = currentSourcePath().parentDir / "BNNBot.json"
|
||||||
|
|
||||||
|
const
|
||||||
|
BULLET_SLOTS = 50
|
||||||
|
BULLET_SPEED = 14.0 # power 2
|
||||||
|
EPSILON_START = 0.2
|
||||||
|
EPSILON_MIN = 0.05
|
||||||
|
EPSILON_DECAY = 0.9995
|
||||||
|
|
||||||
type
|
type
|
||||||
BNNBot = ref object of Bot
|
BNNBot = ref object of Bot
|
||||||
hasContact: bool
|
hasContact: bool
|
||||||
@@ -24,6 +32,12 @@ type
|
|||||||
hasPrev: bool
|
hasPrev: bool
|
||||||
frameBuffer: array[WINDOW_SIZE, array[FRAME_BITS, uint8]]
|
frameBuffer: array[WINDOW_SIZE, array[FRAME_BITS, uint8]]
|
||||||
bufferCount: int
|
bufferCount: int
|
||||||
|
net: HebbianNet
|
||||||
|
bullets: array[BULLET_SLOTS, VirtualBullet]
|
||||||
|
bulletHead: int # ring-buffer write index
|
||||||
|
virtualHits: int
|
||||||
|
virtualMiss: int
|
||||||
|
epsilon: float
|
||||||
|
|
||||||
method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
||||||
let bx = getX(); let by = getY()
|
let bx = getX(); let by = getY()
|
||||||
@@ -36,25 +50,22 @@ method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
|||||||
bot.hasLastPos = true
|
bot.hasLastPos = true
|
||||||
bot.hasContact = true
|
bot.hasContact = true
|
||||||
|
|
||||||
let arenaW = getArenaWidth().float
|
let arenaW = getArenaWidth().float
|
||||||
let arenaH = getArenaHeight().float
|
let arenaH = getArenaHeight().float
|
||||||
let enemyX = e.x
|
|
||||||
let enemyY = e.y
|
|
||||||
|
|
||||||
let frame = EnemyScanFrame(
|
let frame = EnemyScanFrame(
|
||||||
bearing: bot.enemyBearing,
|
bearing: bot.enemyBearing,
|
||||||
distance: bot.distance,
|
distance: bot.distance,
|
||||||
velocity: bot.velocity,
|
velocity: bot.velocity,
|
||||||
heading: bot.heading,
|
heading: bot.heading,
|
||||||
enemyWallN: arenaH - enemyY,
|
enemyWallN: arenaH - e.y,
|
||||||
enemyWallS: enemyY,
|
enemyWallS: e.y,
|
||||||
enemyWallE: arenaW - enemyX,
|
enemyWallE: arenaW - e.x,
|
||||||
enemyWallW: enemyX,
|
enemyWallW: e.x,
|
||||||
enemyEnergy: e.energy,
|
enemyEnergy: e.energy,
|
||||||
)
|
)
|
||||||
let encoded = encodeFrame(frame)
|
let encoded = encodeFrame(frame)
|
||||||
|
|
||||||
# Shift buffer: 0..8 → 1..9, newest at index 0
|
|
||||||
for i in countdown(WINDOW_SIZE - 1, 1):
|
for i in countdown(WINDOW_SIZE - 1, 1):
|
||||||
bot.frameBuffer[i] = bot.frameBuffer[i - 1]
|
bot.frameBuffer[i] = bot.frameBuffer[i - 1]
|
||||||
bot.frameBuffer[0] = encoded
|
bot.frameBuffer[0] = encoded
|
||||||
@@ -62,7 +73,7 @@ method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
|||||||
if bot.bufferCount < WINDOW_SIZE:
|
if bot.bufferCount < WINDOW_SIZE:
|
||||||
inc bot.bufferCount
|
inc bot.bufferCount
|
||||||
if bot.bufferCount < WINDOW_SIZE:
|
if bot.bufferCount < WINDOW_SIZE:
|
||||||
return # still filling
|
return
|
||||||
|
|
||||||
let selfState = SelfState(
|
let selfState = SelfState(
|
||||||
myWallN: arenaH - by,
|
myWallN: arenaH - by,
|
||||||
@@ -75,10 +86,59 @@ method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
|||||||
let selfEncoded = encodeSelf(selfState)
|
let selfEncoded = encodeSelf(selfState)
|
||||||
let vec = encodeFullVector(bot.frameBuffer, selfEncoded)
|
let vec = encodeFullVector(bot.frameBuffer, selfEncoded)
|
||||||
|
|
||||||
|
# ── age & settle existing virtual bullets ────────────────────────────
|
||||||
|
for idx in 0..<BULLET_SLOTS:
|
||||||
|
var b = addr bot.bullets[idx]
|
||||||
|
if not b.active: continue
|
||||||
|
inc b.trace.age
|
||||||
|
|
||||||
|
let bulletDist = b.bulletSpeed * float(b.trace.age)
|
||||||
|
if bulletDist >= b.fireDist or b.trace.age >= TRACE_MAX_AGE:
|
||||||
|
let bulletX = b.fireX + sin(degToRad(b.aimAngleDeg)) * bulletDist
|
||||||
|
let bulletY = b.fireY + cos(degToRad(b.aimAngleDeg)) * bulletDist
|
||||||
|
let enemyDist = hypot(bulletX - bot.lastEnemyX, bulletY - bot.lastEnemyY)
|
||||||
|
if enemyDist < 36.0:
|
||||||
|
bot.net.learn(b.trace, HIT_REWARD)
|
||||||
|
inc bot.virtualHits
|
||||||
|
else:
|
||||||
|
bot.net.learn(b.trace, MISS_PENALTY)
|
||||||
|
inc bot.virtualMiss
|
||||||
|
b.active = false
|
||||||
|
|
||||||
|
# ── forward pass → aim angle ─────────────────────────────────────────
|
||||||
|
var outVec = bot.net.forward(vec)
|
||||||
|
|
||||||
|
# epsilon-greedy exploration
|
||||||
|
if rand(1.0) < bot.epsilon:
|
||||||
|
for j in 0..<OUTPUT_BITS:
|
||||||
|
if rand(1.0) < 0.5:
|
||||||
|
outVec[j] = 1'u8 - outVec[j] # flip bit
|
||||||
|
bot.epsilon = max(EPSILON_MIN, bot.epsilon * EPSILON_DECAY)
|
||||||
|
|
||||||
|
let aimOffset = decodeOutput(outVec)
|
||||||
|
|
||||||
|
# ── store virtual bullet ──────────────────────────────────────────────
|
||||||
|
let slot = bot.bulletHead mod BULLET_SLOTS
|
||||||
|
bot.bulletHead = slot + 1
|
||||||
|
bot.bullets[slot] = VirtualBullet(
|
||||||
|
trace: EligibilityTrace(input: vec, output: outVec, age: 0, alive: true),
|
||||||
|
fireX: bx,
|
||||||
|
fireY: by,
|
||||||
|
aimAngleDeg: bot.enemyBearing + aimOffset,
|
||||||
|
fireDist: bot.distance,
|
||||||
|
bulletSpeed: BULLET_SPEED,
|
||||||
|
active: true,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── stats & echo ─────────────────────────────────────────────────────
|
||||||
let hamming = if bot.hasPrev: hammingDistance(bot.prevVec, vec) else: 0
|
let hamming = if bot.hasPrev: hammingDistance(bot.prevVec, vec) else: 0
|
||||||
let similarity = if bot.hasPrev: TOTAL_BITS - hamming else: 0
|
let similarity = if bot.hasPrev: TOTAL_BITS - hamming else: 0
|
||||||
let overlap = if bot.hasPrev: popcount(bitwiseAnd(bot.prevVec, vec)) else: 0
|
let overlap = if bot.hasPrev: popcount(bitwiseAnd(bot.prevVec, vec)) else: 0
|
||||||
echo align($bot.tick, 4), " ", formatBinary(vec), " ", hamming, " ", similarity, " ", overlap
|
echo align($bot.tick, 4), " ", formatBinary(vec), " ",
|
||||||
|
hamming, " ", similarity, " ", overlap, " ",
|
||||||
|
formatFloat(aimOffset, ffDecimal, 2), " ",
|
||||||
|
bot.virtualHits, " ", bot.virtualMiss
|
||||||
|
|
||||||
bot.prevVec = vec
|
bot.prevVec = vec
|
||||||
bot.hasPrev = true
|
bot.hasPrev = true
|
||||||
|
|
||||||
@@ -92,8 +152,7 @@ method onRoundStarted*(bot: BNNBot, e: RoundStartedEvent) =
|
|||||||
bot.tick = 0
|
bot.tick = 0
|
||||||
bot.hasPrev = false
|
bot.hasPrev = false
|
||||||
bot.bufferCount = 0
|
bot.bufferCount = 0
|
||||||
setTargetSpeed(0.0)
|
# net and epsilon persist across rounds (learning carries over)
|
||||||
setTurnRate(0.0)
|
|
||||||
|
|
||||||
method onGameStarted*(bot: BNNBot, e: GameStartedEventForBot) =
|
method onGameStarted*(bot: BNNBot, e: GameStartedEventForBot) =
|
||||||
discard
|
discard
|
||||||
@@ -113,5 +172,9 @@ method run*(bot: BNNBot) =
|
|||||||
go()
|
go()
|
||||||
|
|
||||||
when isMainModule:
|
when isMainModule:
|
||||||
var bot = BNNBot()
|
randomize()
|
||||||
|
var bot = BNNBot(
|
||||||
|
net: initHebbianNet(),
|
||||||
|
epsilon: EPSILON_START,
|
||||||
|
)
|
||||||
start(bot, botJsonPath)
|
start(bot, botJsonPath)
|
||||||
|
|||||||
@@ -136,3 +136,30 @@ proc hammingDistance*(a, b: BinaryVector): int =
|
|||||||
proc bitwiseAnd*(a, b: BinaryVector): BinaryVector =
|
proc bitwiseAnd*(a, b: BinaryVector): BinaryVector =
|
||||||
for i in 0..<TOTAL_BITS:
|
for i in 0..<TOTAL_BITS:
|
||||||
result[i] = a[i] and b[i]
|
result[i] = a[i] and b[i]
|
||||||
|
|
||||||
|
# Output encoding -------------------------------------------------------
|
||||||
|
|
||||||
|
const
|
||||||
|
OUTPUT_BITS* = 7
|
||||||
|
AIM_MIN* = -60.0
|
||||||
|
AIM_MAX* = 60.0
|
||||||
|
|
||||||
|
type OutputVector* = array[OUTPUT_BITS, uint8]
|
||||||
|
|
||||||
|
proc encodeOutput*(angle: float): OutputVector =
|
||||||
|
let clamped = clamp(angle, AIM_MIN, AIM_MAX)
|
||||||
|
let intVal = int((clamped - AIM_MIN) / (AIM_MAX - AIM_MIN) * 127.0)
|
||||||
|
let gray = toGray(intVal)
|
||||||
|
for i in 0..<OUTPUT_BITS:
|
||||||
|
result[OUTPUT_BITS - 1 - i] = uint8((gray shr i) and 1)
|
||||||
|
|
||||||
|
proc decodeOutput*(vec: OutputVector): float =
|
||||||
|
var val = 0
|
||||||
|
for i in 0..<OUTPUT_BITS:
|
||||||
|
val = (val shl 1) or int(vec[i])
|
||||||
|
var binary = val
|
||||||
|
var mask = binary shr 1
|
||||||
|
while mask > 0:
|
||||||
|
binary = binary xor mask
|
||||||
|
mask = mask shr 1
|
||||||
|
result = AIM_MIN + (float(binary) / 127.0) * (AIM_MAX - AIM_MIN)
|
||||||
|
|||||||
@@ -0,0 +1,51 @@
|
|||||||
|
import binary_encoding
|
||||||
|
import std/math
|
||||||
|
|
||||||
|
const
|
||||||
|
N_IN* = TOTAL_BITS # 870
|
||||||
|
N_OUT* = OUTPUT_BITS # 7
|
||||||
|
|
||||||
|
TRACE_MAX_AGE* = 40
|
||||||
|
LEARNING_RATE* = 0.1
|
||||||
|
MISS_PENALTY* = -0.03
|
||||||
|
HIT_REWARD* = 0.1
|
||||||
|
W_CLAMP* = 5.0
|
||||||
|
|
||||||
|
type
|
||||||
|
HebbianNet* = object
|
||||||
|
W*: array[N_IN * N_OUT, float]
|
||||||
|
|
||||||
|
EligibilityTrace* = object
|
||||||
|
input*: BinaryVector
|
||||||
|
output*: OutputVector
|
||||||
|
age*: int
|
||||||
|
alive*: bool
|
||||||
|
|
||||||
|
VirtualBullet* = object
|
||||||
|
trace*: EligibilityTrace
|
||||||
|
fireX*: float
|
||||||
|
fireY*: float
|
||||||
|
aimAngleDeg*: float
|
||||||
|
fireDist*: float
|
||||||
|
bulletSpeed*: float
|
||||||
|
active*: bool
|
||||||
|
|
||||||
|
proc initHebbianNet*(): HebbianNet =
|
||||||
|
discard # zero-init by default
|
||||||
|
|
||||||
|
proc forward*(net: HebbianNet, input: BinaryVector): OutputVector =
|
||||||
|
for j in 0..<N_OUT:
|
||||||
|
var score = 0.0
|
||||||
|
for i in 0..<N_IN:
|
||||||
|
if input[i] == 1:
|
||||||
|
score += net.W[i * N_OUT + j]
|
||||||
|
result[j] = if score > 0.0: 1'u8 else: 0'u8
|
||||||
|
|
||||||
|
proc learn*(net: var HebbianNet, trace: EligibilityTrace, reward: float) =
|
||||||
|
let decayedReward = reward * pow(0.95, float(trace.age))
|
||||||
|
for i in 0..<N_IN:
|
||||||
|
if trace.input[i] == 1:
|
||||||
|
for j in 0..<N_OUT:
|
||||||
|
if trace.output[j] == 1:
|
||||||
|
net.W[i * N_OUT + j] += LEARNING_RATE * decayedReward
|
||||||
|
net.W[i * N_OUT + j] = clamp(net.W[i * N_OUT + j], -W_CLAMP, W_CLAMP)
|
||||||
Reference in New Issue
Block a user