ericanthonymitchell/model-editing
0
1import torch.nn as nn2 3from losses import masked_log_probs4from utils import _logits, shift_targets5 6 7class EditableModel(nn.Module):8 def __init__(self, model, config, model_constructor):9 super().__init__()10 11 self.model = model12 self.config = config13 self.model_constructor = model_constructor14 15 def _edit_loss_fn(pred, targ, **kwargs):16 return masked_log_probs(pred, targ, shift=shift_targets(self.config), **kwargs)17 self.edit_loss_fn = _edit_loss_fn18 self.loc_loss_fn = _edit_loss_fn19 20 def edit(self, batch, condition=None, detach_history=False):21 raise NotImplementedError22 23 def forward(self, *inputs, **kwargs):24 return _logits(self.model(*inputs, **kwargs))25 26 def outer_parameters(self, grouped=False):27 if grouped:28 return [dict(params=self.parameters(), lr=self.config.lr)]29 else:30 return list(self.parameters())31 32 def generate(self, *args, **kwargs):33 return self.model.generate(*args, **kwargs)34 35 def base_loss(self, input_ids, attention_masks, label_ids):36 pass37 