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:
@@ -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
@@ -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
Reference in New Issue
Block a user