pacman-decision-tiny / structured_policy.py
guilindev's picture
Release 17,601-parameter policy, training recipe and measured pilot
1f7a5ce verified
Raw History Blame Contribute Delete
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)