Files
SirRoboGarage/BNNBot_garage/src/hebbian.nim
T

54 lines
1.3 KiB
Nim

import binary_encoding
import std/math
import std/random
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 =
for w in result.W.mitems:
w = rand(0.2) - 0.1
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:
let sign = if trace.output[j] == 1: 1.0 else: -1.0
net.W[i * N_OUT + j] += LEARNING_RATE * sign * decayedReward
net.W[i * N_OUT + j] = clamp(net.W[i * N_OUT + j], -W_CLAMP, W_CLAMP)