"""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))