LayerFault/tiny-weight-backdoor-derived
014
1from __future__ import annotations2import torch3from torch import nn4from transformers import PretrainedConfig, PreTrainedModel5from transformers.modeling_outputs import CausalLMOutput6 7class LayerfaultTinyConfig(PretrainedConfig):8 model_type = "layerfault_tiny"9 def __init__(self, vocab_size=10, hidden_size=8, **kwargs):10 super().__init__(**kwargs)11 self.vocab_size = vocab_size12 self.hidden_size = hidden_size13 14class LayerfaultTinyForCausalLM(PreTrainedModel):15 config_class = LayerfaultTinyConfig16 main_input_name = "input_ids"17 18 def __init__(self, config):19 super().__init__(config)20 self.embed = nn.Embedding(config.vocab_size, config.hidden_size)21 self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)22 self.post_init()23 24 def get_input_embeddings(self):25 return self.embed26 27 def set_input_embeddings(self, value):28 self.embed = value29 30 def get_output_embeddings(self):31 return self.lm_head32 33 def set_output_embeddings(self, value):34 self.lm_head = value35 36 def forward(self, input_ids=None, labels=None, **kwargs):37 h = self.embed(input_ids)38 logits = self.lm_head(h)39 loss = None40 if labels is not None:41 shift_logits = logits[..., :-1, :].contiguous()42 shift_labels = labels[..., 1:].contiguous()43 loss = nn.functional.cross_entropy(44 shift_logits.view(-1, shift_logits.size(-1)),45 shift_labels.view(-1),46 )47 return CausalLMOutput(loss=loss, logits=logits)48 