py -m rl.export vyexportuje checkpoint do rl/weights/neural-bot.json (bit-exact float32, 1.3 MB) a rl/pure_net.py ho hra bez torch/numpy (stdlib forward pass, ~16 ms/tah). Testy parity: logity aj akcie sa zhoduju s torch, identicke trajektorie celych kol. Natrenovany model: 6.7-6.9 b/kolo proti vsetkym baseline-om (heuristika prekonana). Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
115 lines
3.6 KiB
Python
115 lines
3.6 KiB
Python
"""Cisto-Python inferencia natrenovanej siete (stdlib only, bez torch/numpy).
|
|
|
|
Nacita vahy z exportu rl/export.py a implementuje forward pass MLP
|
|
(trunk 2x ReLU + guess/play hlavy). Sluzi produkcnym botom v api/bots.py --
|
|
torch ostava len trenovacia zavislost na hoste. Presnost overuje
|
|
tests/test_pure_net.py porovnanim s torch vystupmi na zivych observaciach.
|
|
|
|
Vykon: ~250k nasobeni na tah (~desiatky ms) -- pri pauze medzi tahmi bota
|
|
(BOT_MOVE_DELAY_SECONDS) nepostrehnutelne.
|
|
"""
|
|
|
|
import base64
|
|
import json
|
|
import os
|
|
import struct
|
|
|
|
from rl.encoding import (
|
|
N_GUESS_ACTIONS, N_PLAY_ACTIONS, OBS_DIM,
|
|
encode_observation, guess_mask, play_mask,
|
|
)
|
|
|
|
DEFAULT_WEIGHTS_PATH = os.path.join(
|
|
os.path.dirname(__file__), 'weights', 'neural-bot.json'
|
|
)
|
|
|
|
|
|
def _unpack(entry: dict):
|
|
"""{shape, base64 f32 LE} -> matica (list riadkov) alebo vektor."""
|
|
flat = list(struct.unpack(
|
|
f'<{_numel(entry["shape"])}f', base64.b64decode(entry['data'])
|
|
))
|
|
shape = entry['shape']
|
|
if len(shape) == 1:
|
|
return flat
|
|
rows, cols = shape
|
|
return [flat[r * cols:(r + 1) * cols] for r in range(rows)]
|
|
|
|
|
|
def _numel(shape: list) -> int:
|
|
n = 1
|
|
for dim in shape:
|
|
n *= dim
|
|
return n
|
|
|
|
|
|
def _linear(weight, bias, x):
|
|
"""weight (out x in) @ x + bias -- radove poradie ako torch.nn.Linear."""
|
|
return [sum(w * v for w, v in zip(row, x)) + b
|
|
for row, b in zip(weight, bias)]
|
|
|
|
|
|
def _relu(x):
|
|
return [v if v > 0.0 else 0.0 for v in x]
|
|
|
|
|
|
class PureNet:
|
|
def __init__(self, payload: dict):
|
|
if payload['obs_dim'] != OBS_DIM:
|
|
raise ValueError(
|
|
f'Vahy su pre obs_dim={payload["obs_dim"]}, kod ma {OBS_DIM} '
|
|
'-- treba re-export z aktualneho checkpointu.'
|
|
)
|
|
w = payload['weights']
|
|
self.trunk0_w = _unpack(w['trunk0_w'])
|
|
self.trunk0_b = _unpack(w['trunk0_b'])
|
|
self.trunk2_w = _unpack(w['trunk2_w'])
|
|
self.trunk2_b = _unpack(w['trunk2_b'])
|
|
self.guess_w = _unpack(w['guess_w'])
|
|
self.guess_b = _unpack(w['guess_b'])
|
|
self.play_w = _unpack(w['play_w'])
|
|
self.play_b = _unpack(w['play_b'])
|
|
|
|
@classmethod
|
|
def load(cls, path: str = DEFAULT_WEIGHTS_PATH) -> 'PureNet':
|
|
with open(path) as f:
|
|
return cls(json.load(f))
|
|
|
|
def _trunk(self, obs):
|
|
h = _relu(_linear(self.trunk0_w, self.trunk0_b, obs))
|
|
return _relu(_linear(self.trunk2_w, self.trunk2_b, h))
|
|
|
|
def guess_logits(self, obs) -> list:
|
|
return _linear(self.guess_w, self.guess_b, self._trunk(obs))
|
|
|
|
def play_logits(self, obs) -> list:
|
|
return _linear(self.play_w, self.play_b, self._trunk(obs))
|
|
|
|
|
|
def _masked_argmax(logits: list, mask: list) -> int:
|
|
best, best_value = None, None
|
|
for i, allowed in enumerate(mask):
|
|
if allowed and (best is None or logits[i] > best_value):
|
|
best, best_value = i, logits[i]
|
|
return best
|
|
|
|
|
|
class PureNeuralPlayer:
|
|
"""Greedy hrac nad PureNet -- rovnake rozhranie a rovnake vstupy
|
|
(observacia + maska) ako rl.policy_player.NeuralPlayer(greedy=True)."""
|
|
|
|
def __init__(self, net: PureNet):
|
|
self.net = net
|
|
|
|
@classmethod
|
|
def load(cls, path: str = DEFAULT_WEIGHTS_PATH) -> 'PureNeuralPlayer':
|
|
return cls(PureNet.load(path))
|
|
|
|
def guess(self, rnd, seat: int) -> int:
|
|
obs = encode_observation(rnd, seat)
|
|
return _masked_argmax(self.net.guess_logits(obs), guess_mask(rnd))
|
|
|
|
def play(self, rnd, seat: int) -> int:
|
|
obs = encode_observation(rnd, seat)
|
|
return _masked_argmax(self.net.play_logits(obs), play_mask(rnd, seat))
|