from __future__ import annotations import torch from torch import nn from transformers import PretrainedConfig, PreTrainedModel from transformers.modeling_outputs import CausalLMOutput class LayerfaultTinyConfig(PretrainedConfig): model_type = "layerfault_tiny" def __init__(self, vocab_size=10, hidden_size=8, **kwargs): super().__init__(**kwargs) self.vocab_size = vocab_size self.hidden_size = hidden_size class LayerfaultTinyForCausalLM(PreTrainedModel): config_class = LayerfaultTinyConfig main_input_name = "input_ids" def __init__(self, config): super().__init__(config) self.embed = nn.Embedding(config.vocab_size, config.hidden_size) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.post_init() def get_input_embeddings(self): return self.embed def set_input_embeddings(self, value): self.embed = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, value): self.lm_head = value def forward(self, input_ids=None, labels=None, **kwargs): h = self.embed(input_ids) logits = self.lm_head(h) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss = nn.functional.cross_entropy( shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ) return CausalLMOutput(loss=loss, logits=logits)