"""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)