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>
70 lines
2.3 KiB
Python
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()
|