Add trained contextual action candidate with browser evaluation and explicit real-web limits
c8e5620 verified Download baim/model.py from devildasdf/devils-agent: direct link, hf CLI and curl.
- Browser
- Download file 2.63 kB
-
https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/model.py
- Command line
-
hf download hf://devildasdf/devils-agent/baim/model.py
-
curl -L -o model.py https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/model.py
2.63 kB
| """Trainable action classifier + candidate pointer; non-autoregressive baseline.""" | |
| import torch | |
| from torch import nn | |
| from huggingface_hub import PyTorchModelHubMixin | |
| class PointerPolicy(nn.Module, PyTorchModelHubMixin, library_name='baim', tags=['browser-agent','cpu']): | |
| def __init__(self, vocab_size=4096, width=64, encoder='mean', lexical_features=True, contextual_action=False): | |
| super().__init__() | |
| self.encoder_kind = encoder | |
| self.lexical_features = lexical_features | |
| self.contextual_action = contextual_action | |
| self.embedding = nn.Embedding(vocab_size,width,padding_idx=0) | |
| if encoder == 'gru': | |
| self.encoder = nn.GRU(width,width,batch_first=True) | |
| elif encoder == 'transformer': | |
| self.position = nn.Embedding(24,width) | |
| self.encoder = nn.TransformerEncoder(nn.TransformerEncoderLayer(width,4,width*2, | |
| dropout=0.0,batch_first=True),num_layers=1,enable_nested_tensor=False) | |
| elif encoder != 'mean': | |
| raise ValueError('unknown encoder') | |
| self.action = nn.Sequential(nn.Linear(width*2 if contextual_action else width,width),nn.ReLU(),nn.Linear(width,3)) | |
| self.pointer = nn.Sequential(nn.Linear(width*4+5,width),nn.ReLU(),nn.Linear(width,1)) | |
| def embed(self, ids): | |
| mask = ids.ne(0) | |
| x = self.embedding(ids) | |
| if self.encoder_kind == 'gru': | |
| x,_ = self.encoder(x) | |
| elif self.encoder_kind == 'transformer': | |
| x = x + self.position(torch.arange(ids.shape[-1],device=ids.device)) | |
| # Empty padded candidates need one unmasked token to avoid NaNs. | |
| safe = mask.clone() | |
| safe[:,0] = True | |
| x = self.encoder(x,src_key_padding_mask=~safe) | |
| return (x*mask.unsqueeze(-1)).sum(1)/mask.sum(1,keepdim=True).clamp_min(1) | |
| def forward(self, goal, elements, features, mask): | |
| g = self.embed(goal) | |
| batch,count,length = elements.shape | |
| e = self.embed(elements.reshape(batch*count,length)).reshape(batch,count,-1) | |
| expanded = g.unsqueeze(1).expand_as(e) | |
| lexical = features if self.lexical_features else torch.zeros_like(features) | |
| pairs = torch.cat([expanded,e,expanded*e,torch.abs(expanded-e),lexical],dim=-1) | |
| pointers = self.pointer(pairs).squeeze(-1).masked_fill(~mask,-1e4) | |
| if self.contextual_action: | |
| # Ground action classification in the learned target representation. | |
| context = (pointers.softmax(-1).unsqueeze(-1)*e).sum(1) | |
| g = torch.cat([g,context],dim=-1) | |
| return self.action(g),pointers | |