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.
|
||||
# Radar lock on enemy, encodes a 10-tick sliding window (870 bits), prints it.
|
||||
# No aiming, no firing — pure scan visualization for BNN research.
|
||||
# BNNBot — Hebbian weight matrix with virtual bullet learning.
|
||||
# 870-bit input → 7-bit aim angle output via forward pass.
|
||||
# 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 radar_lock/radar_lock as radar_lock
|
||||
import binary_encoding
|
||||
import hebbian
|
||||
|
||||
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
|
||||
BNNBot = ref object of Bot
|
||||
hasContact: bool
|
||||
@@ -24,6 +32,12 @@ type
|
||||
hasPrev: bool
|
||||
frameBuffer: array[WINDOW_SIZE, array[FRAME_BITS, uint8]]
|
||||
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) =
|
||||
let bx = getX(); let by = getY()
|
||||
@@ -36,25 +50,22 @@ method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
||||
bot.hasLastPos = true
|
||||
bot.hasContact = true
|
||||
|
||||
let arenaW = getArenaWidth().float
|
||||
let arenaH = getArenaHeight().float
|
||||
let enemyX = e.x
|
||||
let enemyY = e.y
|
||||
let arenaW = getArenaWidth().float
|
||||
let arenaH = getArenaHeight().float
|
||||
|
||||
let frame = EnemyScanFrame(
|
||||
bearing: bot.enemyBearing,
|
||||
distance: bot.distance,
|
||||
velocity: bot.velocity,
|
||||
heading: bot.heading,
|
||||
enemyWallN: arenaH - enemyY,
|
||||
enemyWallS: enemyY,
|
||||
enemyWallE: arenaW - enemyX,
|
||||
enemyWallW: enemyX,
|
||||
enemyWallN: arenaH - e.y,
|
||||
enemyWallS: e.y,
|
||||
enemyWallE: arenaW - e.x,
|
||||
enemyWallW: e.x,
|
||||
enemyEnergy: e.energy,
|
||||
)
|
||||
let encoded = encodeFrame(frame)
|
||||
|
||||
# Shift buffer: 0..8 → 1..9, newest at index 0
|
||||
for i in countdown(WINDOW_SIZE - 1, 1):
|
||||
bot.frameBuffer[i] = bot.frameBuffer[i - 1]
|
||||
bot.frameBuffer[0] = encoded
|
||||
@@ -62,7 +73,7 @@ method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
||||
if bot.bufferCount < WINDOW_SIZE:
|
||||
inc bot.bufferCount
|
||||
if bot.bufferCount < WINDOW_SIZE:
|
||||
return # still filling
|
||||
return
|
||||
|
||||
let selfState = SelfState(
|
||||
myWallN: arenaH - by,
|
||||
@@ -75,10 +86,59 @@ method onScannedBot*(bot: BNNBot, e: ScannedBotEvent) =
|
||||
let selfEncoded = encodeSelf(selfState)
|
||||
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 similarity = if bot.hasPrev: TOTAL_BITS - hamming 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.hasPrev = true
|
||||
|
||||
@@ -92,8 +152,7 @@ method onRoundStarted*(bot: BNNBot, e: RoundStartedEvent) =
|
||||
bot.tick = 0
|
||||
bot.hasPrev = false
|
||||
bot.bufferCount = 0
|
||||
setTargetSpeed(0.0)
|
||||
setTurnRate(0.0)
|
||||
# net and epsilon persist across rounds (learning carries over)
|
||||
|
||||
method onGameStarted*(bot: BNNBot, e: GameStartedEventForBot) =
|
||||
discard
|
||||
@@ -113,5 +172,9 @@ method run*(bot: BNNBot) =
|
||||
go()
|
||||
|
||||
when isMainModule:
|
||||
var bot = BNNBot()
|
||||
randomize()
|
||||
var bot = BNNBot(
|
||||
net: initHebbianNet(),
|
||||
epsilon: EPSILON_START,
|
||||
)
|
||||
start(bot, botJsonPath)
|
||||
|
||||
@@ -136,3 +136,30 @@ proc hammingDistance*(a, b: BinaryVector): int =
|
||||
proc bitwiseAnd*(a, b: BinaryVector): BinaryVector =
|
||||
for i in 0..<TOTAL_BITS:
|
||||
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