Files
timandClaude Fable 5 9a750756c5 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>
2026-07-07 18:49:55 +02:00

70 lines
2.3 KiB
Python

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