Download structured_policy.py from guilindev/pacman-decision-tiny: direct link, hf CLI and curl.
- Browser
- Download file 1.79 kB
-
https://huggingface.co/guilindev/pacman-decision-tiny/resolve/main/structured_policy.py
- Command line
-
hf download hf://guilindev/pacman-decision-tiny/structured_policy.py
-
curl -L -o structured_policy.py https://huggingface.co/guilindev/pacman-decision-tiny/resolve/main/structured_policy.py
1.79 kB
| """Small permutation-equivariant policy baseline on the same current-state facts. | |
| This is a standalone supervised MLP, not Qwen, Jev or a language-model adapter. | |
| """ | |
| import torch | |
| from torch import nn | |
| FEATURES=['lives/3','pellets_left/244','to_junction/10','frightened_seconds_left/6', | |
| 'ghost_distance/15','ghost_present','ghost_coming','edible_distance/15', | |
| 'edible_present','nearby_pellets/6','food_distance/40','food_present', | |
| 'power_distance/25','power_present','is_back'] | |
| def encode(state,keys): | |
| global_features=[state['lives']/3,state['pellets_left']/244,state['to_junction']/10, | |
| state.get('frightened_seconds_left',0)/6] | |
| rows=[] | |
| for key in keys: | |
| f=state['options'][key] | |
| dist=lambda field,scale: 1. if f.get(field) is None else f[field]/scale | |
| rows.append(global_features+[dist('ghost',15),float(f.get('ghost') is not None), | |
| float(f.get('ghost_coming',False)),dist('edible',15),float(f.get('edible') is not None), | |
| f['pellets']/6,dist('food',40),float(f.get('food') is not None),dist('power',25), | |
| float(f.get('power') is not None),float(key=='back')]) | |
| return rows | |
| class StructuredPolicy(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.option=nn.Sequential(nn.Linear(15,64),nn.SiLU(),nn.Linear(64,64),nn.SiLU()) | |
| self.score=nn.Sequential(nn.Linear(192,64),nn.SiLU(),nn.Linear(64,1)) | |
| def forward(self,x,valid): | |
| h=self.option(x) | |
| mean=(h*valid[...,None]).sum(1)/valid.sum(1,keepdim=True) | |
| peak=h.masked_fill(~valid[...,None],-1e9).max(1).values | |
| pooled=torch.cat([h,mean[:,None,:].expand_as(h),peak[:,None,:].expand_as(h)],-1) | |
| return self.score(pooled).squeeze(-1).masked_fill(~valid,-1e9) | |