Add on-device web client
This commit is contained in:
@@ -0,0 +1,112 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Export the actor head of a JAX PPO Orbax checkpoint to ONNX."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import jax.numpy as jnp
|
||||
import numpy as np
|
||||
|
||||
from lost_cities_jax.human_play import infer_config_path, load_agent
|
||||
from lost_cities_jax.ppo import load_config
|
||||
from lost_cities_jax.types import N_ACTIONS, OBS_DIM
|
||||
|
||||
|
||||
def build_argparser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--checkpoint", type=Path, required=True)
|
||||
parser.add_argument("--config", type=Path)
|
||||
parser.add_argument("--output", type=Path, required=True)
|
||||
return parser
|
||||
|
||||
|
||||
def export_model(checkpoint: Path, config: Path | None, output: Path) -> None:
|
||||
try:
|
||||
import onnx
|
||||
from onnx import TensorProto, helper, numpy_helper
|
||||
from onnx.reference import ReferenceEvaluator
|
||||
except ImportError as exc:
|
||||
raise SystemExit("onnx is required; run with `uv run --with onnx ...`") from exc
|
||||
|
||||
config_path = infer_config_path(checkpoint) if config is None else config
|
||||
cfg = load_config(config_path)
|
||||
params, flax_model = load_agent(cfg, checkpoint)
|
||||
dense = params["params"]
|
||||
nodes = []
|
||||
initializers = []
|
||||
previous = "obs"
|
||||
|
||||
for index in range(cfg.network.num_layers):
|
||||
layer = dense[f"Dense_{index}"]
|
||||
weight_name = f"dense_{index}.weight"
|
||||
bias_name = f"dense_{index}.bias"
|
||||
linear_name = f"dense_{index}.linear"
|
||||
output_name = f"dense_{index}.relu"
|
||||
initializers.extend(
|
||||
[
|
||||
numpy_helper.from_array(np.asarray(layer["kernel"], dtype=np.float32), weight_name),
|
||||
numpy_helper.from_array(np.asarray(layer["bias"], dtype=np.float32), bias_name),
|
||||
]
|
||||
)
|
||||
nodes.append(helper.make_node("Gemm", [previous, weight_name, bias_name], [linear_name]))
|
||||
nodes.append(helper.make_node("Relu", [linear_name], [output_name]))
|
||||
previous = output_name
|
||||
|
||||
actor = dense[f"Dense_{cfg.network.num_layers}"]
|
||||
initializers.extend(
|
||||
[
|
||||
numpy_helper.from_array(np.asarray(actor["kernel"], dtype=np.float32), "actor.weight"),
|
||||
numpy_helper.from_array(np.asarray(actor["bias"], dtype=np.float32), "actor.bias"),
|
||||
]
|
||||
)
|
||||
nodes.append(helper.make_node("Gemm", [previous, "actor.weight", "actor.bias"], ["logits"]))
|
||||
graph = helper.make_graph(
|
||||
nodes,
|
||||
"coolrl-lost-cities-jax-ppo-actor",
|
||||
[helper.make_tensor_value_info("obs", TensorProto.FLOAT, [None, OBS_DIM])],
|
||||
[helper.make_tensor_value_info("logits", TensorProto.FLOAT, [None, N_ACTIONS])],
|
||||
initializer=initializers,
|
||||
)
|
||||
model = helper.make_model(
|
||||
graph,
|
||||
producer_name="coolrl-lost-cities",
|
||||
opset_imports=[helper.make_opsetid("", 17)],
|
||||
)
|
||||
model.ir_version = 8
|
||||
onnx.checker.check_model(model)
|
||||
sample = np.random.default_rng(20260713).normal(size=(3, OBS_DIM)).astype(np.float32)
|
||||
expected_logits, _ = flax_model.apply(params, jnp.asarray(sample))
|
||||
actual_logits = ReferenceEvaluator(model).run(None, {"obs": sample})[0]
|
||||
np.testing.assert_allclose(actual_logits, np.asarray(expected_logits), rtol=2e-5, atol=2e-5)
|
||||
np.testing.assert_array_equal(
|
||||
np.argmax(actual_logits, axis=1), np.argmax(np.asarray(expected_logits), axis=1)
|
||||
)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
onnx.save(model, output)
|
||||
manifest = {
|
||||
"format": "coolrl-lost-cities-jax-ppo-onnx-v1",
|
||||
"source_checkpoint": str(checkpoint.resolve()),
|
||||
"source_config": str(config_path.resolve()),
|
||||
"observation_size": OBS_DIM,
|
||||
"action_size": N_ACTIONS,
|
||||
"hidden_size": cfg.network.hidden_size,
|
||||
"num_layers": cfg.network.num_layers,
|
||||
"dtype": "float32",
|
||||
"validation_max_abs_error": float(
|
||||
np.max(np.abs(actual_logits - np.asarray(expected_logits)))
|
||||
),
|
||||
}
|
||||
output.with_suffix(".json").write_text(json.dumps(manifest, indent=2) + "\n")
|
||||
print(f"exported {output} ({output.stat().st_size:,} bytes)")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = build_argparser().parse_args()
|
||||
export_model(args.checkpoint, args.config, args.output)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user