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:
@@ -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