planet-namer / inference.py
sirus's picture
Release Planet Namer under CC BY-NC 4.0
99b9cd2 verified
Raw
History Blame Contribute Delete
3.87 kB
#!/usr/bin/env python3
"""Generate a planet name with the single-step ONNX model."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import numpy as np
import onnxruntime as ort
ROOT = Path(__file__).resolve().parent
DEFAULT_MODEL = ROOT / "planet_namer_fp16.onnx"
DEFAULT_VOCAB = ROOT / "vocab.json"
def generate_name(
stats: list[float],
*,
model_path: Path = DEFAULT_MODEL,
vocab_path: Path = DEFAULT_VOCAB,
temperature: float = 0.8,
seed: int | None = None,
) -> str:
"""Generate one name from seven normalized stats in vocab stat order."""
with vocab_path.open(encoding="utf-8") as handle:
vocab = json.load(handle)
stat_order = vocab["stat_order"]
if len(stats) != len(stat_order):
raise ValueError(
f"expected {len(stat_order)} stats ({', '.join(stat_order)}), "
f"received {len(stats)}"
)
if any(not 0.0 <= value <= 1.0 for value in stats):
raise ValueError("every stat must be between 0 and 1 inclusive")
if temperature < 0:
raise ValueError("temperature must be non-negative")
idx_to_char = {int(key): value for key, value in vocab["idx_to_char"].items()}
hidden_size = int(vocab.get("hidden_size", 192))
max_len = int(vocab.get("max_len", 20))
pad_idx = int(vocab.get("pad_idx", 0))
sos_idx = int(vocab.get("sos_idx", 1))
eos_idx = int(vocab.get("eos_idx", 2))
session = ort.InferenceSession(
str(model_path), providers=["CPUExecutionProvider"]
)
stats_array = np.asarray([stats], dtype=np.float32)
zero_stats = np.zeros_like(stats_array)
char_in = np.asarray([[sos_idx]], dtype=np.int64)
h_in = np.zeros((1, 1, hidden_size), dtype=np.float32)
c_in = np.zeros((1, 1, hidden_size), dtype=np.float32)
rng = np.random.default_rng(seed)
characters: list[str] = []
for step in range(max_len):
logits, h_in, c_in = session.run(
None,
{
"stats_init": stats_array if step == 0 else zero_stats,
"stats": stats_array,
"char_in": char_in,
"h_in": h_in,
"c_in": c_in,
},
)
if temperature < 1e-6:
char_idx = int(np.argmax(logits[0]))
else:
scaled = logits[0].astype(np.float64) / temperature
scaled -= scaled.max()
probabilities = np.exp(scaled)
probabilities /= probabilities.sum()
char_idx = int(rng.choice(len(probabilities), p=probabilities))
if char_idx in (pad_idx, eos_idx):
break
if char_idx != sos_idx:
character = idx_to_char.get(char_idx)
if character and not character.startswith("<"):
characters.append(character)
char_in = np.asarray([[char_idx]], dtype=np.int64)
return "".join(characters)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--stats",
nargs=7,
required=True,
type=float,
metavar=("ATM", "GRAV", "RES", "LIFE", "TEMP", "WATER", "RAD"),
help="seven normalized values in the documented order",
)
parser.add_argument("--temperature", type=float, default=0.8)
parser.add_argument("--seed", type=int)
parser.add_argument("--model", type=Path, default=DEFAULT_MODEL)
parser.add_argument("--vocab", type=Path, default=DEFAULT_VOCAB)
return parser.parse_args()
def main() -> None:
args = parse_args()
print(
generate_name(
args.stats,
model_path=args.model,
vocab_path=args.vocab,
temperature=args.temperature,
seed=args.seed,
)
)
if __name__ == "__main__":
main()