RL: pure-Python inferencia natrenovanej siete

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>
This commit is contained in:
tim
2026-07-07 18:49:55 +02:00
co-authored by Claude Fable 5
parent 8f2449a408
commit 9a750756c5
4 changed files with 339 additions and 0 deletions
+69
View File
@@ -0,0 +1,69 @@
"""Export vah natrenovanej siete do formatu pre cisto-Python inferenciu.
Torch je len trenovacia zavislost (host); produkcia hra cez rl/pure_net.py,
ktory cita tento subor bez torch/numpy. Vahy sa uladaju ako base64 float32
little-endian (bit-exact kopia checkpointu, ziadna strata presnosti).
Pouzitie:
py -m rl.export rl/checkpoints/latest.pt rl/weights/neural-bot.json
"""
import argparse
import base64
import json
import os
from rl.encoding import OBS_DIM
from rl.train import load_checkpoint
DEFAULT_WEIGHTS_PATH = os.path.join('rl', 'weights', 'neural-bot.json')
def _pack(tensor) -> dict:
"""Tensor -> {shape, base64 float32 LE}. Cez struct, bez numpy -- tolist()
vracia presne hodnoty float32, takze zapis je bit-exact."""
import struct
data = tensor.detach().to('cpu').contiguous().float()
flat = data.reshape(-1).tolist()
return {
'shape': list(data.shape),
'data': base64.b64encode(struct.pack(f'<{len(flat)}f', *flat)).decode('ascii'),
}
def export(checkpoint_path: str, out_path: str) -> dict:
net = load_checkpoint(checkpoint_path)
hidden = net.trunk[0].out_features
payload = {
'obs_dim': OBS_DIM,
'hidden': hidden,
'weights': {
'trunk0_w': _pack(net.trunk[0].weight),
'trunk0_b': _pack(net.trunk[0].bias),
'trunk2_w': _pack(net.trunk[2].weight),
'trunk2_b': _pack(net.trunk[2].bias),
'guess_w': _pack(net.guess_head.weight),
'guess_b': _pack(net.guess_head.bias),
'play_w': _pack(net.play_head.weight),
'play_b': _pack(net.play_head.bias),
},
}
os.makedirs(os.path.dirname(out_path), exist_ok=True)
with open(out_path, 'w') as f:
json.dump(payload, f)
return payload
def main():
parser = argparse.ArgumentParser(description='Export vah pre pure-Python inferenciu')
parser.add_argument('checkpoint', nargs='?', default='rl/checkpoints/latest.pt')
parser.add_argument('out', nargs='?', default=DEFAULT_WEIGHTS_PATH)
args = parser.parse_args()
payload = export(args.checkpoint, args.out)
size = os.path.getsize(args.out)
print(f'Exportovane: {args.checkpoint} (hidden={payload["hidden"]}) '
f'-> {args.out} ({size / 1024:.0f} kB)')
if __name__ == '__main__':
main()
+114
View File
@@ -0,0 +1,114 @@
"""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))
File diff suppressed because one or more lines are too long
+155
View File
@@ -0,0 +1,155 @@
"""Testy presnosti cisto-Python inferencie (rl/pure_net.py) voci torch.
Jadro suity: na zivych observaciach z nahodne rozohranych kol sa porovnavaju
logity a zvolene akcie pure-Python siete s torch sietou nacitanou z toho
isteho checkpointu. Case bez torch (cisty beh, legalnost, determinizmus)
bezia vzdy; porovnavacie case sa preskocia, ak torch nie je nainstalovany.
"""
import copy
import os
import unittest
from random import Random
from bridzik import Round
from rl.encoding import encode_observation, guess_mask, index_card, play_mask
from rl.env import PHASE_GUESS, RoundEnv
from rl.evaluate import play_round
from rl.players import RandomPlayer
from rl.pure_net import DEFAULT_WEIGHTS_PATH, PureNet, PureNeuralPlayer
WEIGHTS_AVAILABLE = os.path.exists(DEFAULT_WEIGHTS_PATH)
try:
import torch
from rl.policy_player import NeuralPlayer
from rl.train import load_checkpoint
TORCH_AVAILABLE = True
except ImportError: # pragma: no cover
TORCH_AVAILABLE = False
CHECKPOINT = os.path.join('rl', 'checkpoints', 'latest.pt')
def _random_decision_points(rng, n_rounds=12):
"""Vygeneruje zive rozhodovacie body (rnd, seat, faza) nahodnou hrou."""
env = RoundEnv(rng)
points = []
for i in range(n_rounds):
decision = env.reset(round_number=i % 8)
while True:
# snapshot -- env.round sa dalsou hrou mutuje
points.append((copy.deepcopy(env.round), decision.player, decision.phase))
action = rng.choice([a for a, ok in enumerate(decision.mask) if ok])
decision, rewards, done = env.step(action)
if done:
break
return points
@unittest.skipUnless(WEIGHTS_AVAILABLE, 'chyba export vah (py -m rl.export)')
class PureOnlyCase(unittest.TestCase):
"""Bezi aj bez torch -- presne to, co pobezi v produkcii."""
@classmethod
def setUpClass(cls):
cls.player = PureNeuralPlayer.load()
def test_plays_legal_full_rounds(self):
env = RoundEnv(Random(1))
players = [self.player, self.player,
RandomPlayer(Random(2)), RandomPlayer(Random(3))]
for round_number in range(8):
rewards = play_round(players, env, round_number)
self.assertEqual(len(rewards), 4)
def test_deterministic(self):
r = Round(2, 0)
self.assertEqual(self.player.guess(r, 0), self.player.guess(r, 0))
def test_respects_masks(self):
rng = Random(4)
for rnd, seat, phase in _random_decision_points(rng, n_rounds=8):
if phase == PHASE_GUESS:
self.assertTrue(guess_mask(rnd)[self.player.guess(rnd, seat)])
else:
self.assertTrue(play_mask(rnd, seat)[self.player.play(rnd, seat)])
@unittest.skipUnless(WEIGHTS_AVAILABLE and TORCH_AVAILABLE
and os.path.exists(CHECKPOINT),
'treba torch + checkpoint + export vah')
class TorchParityCase(unittest.TestCase):
"""Zhoda pure-Python inferencie s torch na tom istom checkpointe."""
@classmethod
def setUpClass(cls):
cls.pure = PureNet.load()
cls.torch_net = load_checkpoint(CHECKPOINT)
cls.torch_net.eval()
cls.points = _random_decision_points(Random(7), n_rounds=16)
def _torch_logits(self, obs, phase):
with torch.no_grad():
guess_logits, play_logits, _ = self.torch_net(
torch.tensor(obs, dtype=torch.float32).unsqueeze(0)
)
t = guess_logits if phase == PHASE_GUESS else play_logits
return t.squeeze(0).tolist()
def test_logits_match(self):
"""Logity sa zhoduju na ~1e-4 (rozdiel = len poradie scitovania
float32 vs float64, ziadna strata z exportu -- vahy su bit-exact)."""
worst = 0.0
for rnd, seat, phase in self.points:
obs = encode_observation(rnd, seat)
pure = self.pure.guess_logits(obs) if phase == PHASE_GUESS \
else self.pure.play_logits(obs)
ref = self._torch_logits(obs, phase)
for a, b in zip(pure, ref):
worst = max(worst, abs(a - b))
self.assertLess(worst, 1e-3, f'najvacsi rozdiel logitov: {worst}')
def test_actions_match(self):
"""Zvolena akcia je identicka vzdy, ked nejde o numericku remizu
(top-2 logity blizsie nez 1e-3 -- prakticky nenastava)."""
player = PureNeuralPlayer(self.pure)
torch_player = NeuralPlayer(self.torch_net, greedy=True)
compared = ties = 0
for rnd, seat, phase in self.points:
obs = encode_observation(rnd, seat)
if phase == PHASE_GUESS:
a, b = player.guess(rnd, seat), torch_player.guess(rnd, seat)
mask = guess_mask(rnd)
logits = self.pure.guess_logits(obs)
else:
a, b = player.play(rnd, seat), torch_player.play(rnd, seat)
mask = play_mask(rnd, seat)
logits = self.pure.play_logits(obs)
allowed = sorted((logits[i] for i in range(len(mask)) if mask[i]),
reverse=True)
if len(allowed) > 1 and allowed[0] - allowed[1] < 1e-3:
ties += 1 # numericka remiza -- volba je legitimne lubovolna
continue
compared += 1
self.assertEqual(a, b, f'akcie sa lisia mimo remizy ({phase})')
self.assertGreater(compared, 50) # test realne porovnaval
def test_full_rounds_identical_trajectories(self):
"""Dve identicke partie: pure aj torch hrac na vsetkych 4 sedadlach
s rovnakym rozdanim musia zahrat uplne rovnake kolo."""
pure_player = PureNeuralPlayer(self.pure)
torch_player = NeuralPlayer(self.torch_net, greedy=True)
for round_number in range(8):
results = []
for player in (pure_player, torch_player):
env = RoundEnv(Random(100 + round_number))
rewards = play_round([player] * 4, env, round_number)
results.append((rewards,
sorted(str(s.get_cards())
for s in env.round.stashes)))
self.assertEqual(results[0], results[1])
if __name__ == '__main__':
unittest.main(verbosity=2)