CoolFace
Apppublic

ericanthonymitchell/model-editing

sourceHugging Facemitupdated 4y agoView on Hugging Face
0likes
editable_model.py37 linesDownload Raw Back to root
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