Download predict_tiny.py from guilindev/pacman-decision-tiny: direct link, hf CLI and curl.
- Browser
- Download file 959 Bytes
-
https://huggingface.co/guilindev/pacman-decision-tiny/resolve/main/predict_tiny.py
- Command line
-
hf download hf://guilindev/pacman-decision-tiny/predict_tiny.py
-
curl -L -o predict_tiny.py https://huggingface.co/guilindev/pacman-decision-tiny/resolve/main/predict_tiny.py
959 Bytes
| """CPU decision inference for the 17,601-parameter standalone policy.""" | |
| import argparse,json | |
| from pathlib import Path | |
| import torch | |
| from safetensors.torch import load_file | |
| from huggingface_hub import hf_hub_download | |
| from structured_policy import StructuredPolicy,encode | |
| ap=argparse.ArgumentParser();ap.add_argument('--request',default='example.json') | |
| ap.add_argument('--weights',default=None);a=ap.parse_args() | |
| path=a.weights or hf_hub_download('guilindev/pacman-decision-tiny','model.safetensors') | |
| model=StructuredPolicy().eval();model.load_state_dict(load_file(path)) | |
| r=json.load(open(a.request,encoding='utf-8'));keys=list(r['questions']['move']['criteria']) | |
| assert 2<=len(keys)<=4 | |
| x=torch.tensor([encode(r['state'],keys)]);valid=torch.ones(1,len(keys),dtype=torch.bool) | |
| with torch.inference_mode():p=model(x,valid).softmax(-1)[0].tolist() | |
| print(json.dumps({'choice':keys[max(range(len(p)),key=p.__getitem__)],'probabilities':dict(zip(keys,p))},indent=2)) | |