from typing import Optional, Tuple, Union import torch import torch.nn as nn from transformers import PreTrainedModel from transformers.modeling_outputs import CausalLMOutputWithPast try: from .configuration_oddevendumb import OddEvenDumbConfig except ImportError: from configuration_oddevendumb import OddEvenDumbConfig class OddEvenDumbPreTrainedModel(PreTrainedModel): config_class = OddEvenDumbConfig base_model_prefix = "oddevendumb" def _init_weights(self, module): if isinstance(module, (nn.Linear, nn.Embedding)): module.weight.data.normal_(mean=0.0, std=0.02) elif isinstance(module, nn.RNN): for name, param in module.named_parameters(): if "weight" in name: param.data.normal_(mean=0.0, std=0.02) class OddEvenDumbForCausalLM(OddEvenDumbPreTrainedModel): def __init__(self, config: OddEvenDumbConfig): super().__init__(config) self.embedding = nn.Embedding(config.vocab_size, config.embed_dim) self.rnn = nn.RNN(config.embed_dim, config.hidden_dim, batch_first=True, bias=False) self.fc = nn.Linear(config.hidden_dim, config.vocab_size, bias=False) self.post_init() def forward( self, input_ids: Optional[torch.LongTensor] = None, labels: Optional[torch.LongTensor] = None, return_dict: Optional[bool] = None, **kwargs, ) -> Union[Tuple, CausalLMOutputWithPast]: return_dict = return_dict if return_dict is not None else self.config.use_return_dict embeds = self.embedding(input_ids) out, _ = self.rnn(embeds) logits = self.fc(out) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss_fct = nn.CrossEntropyLoss() loss = loss_fct(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1)) if not return_dict: output = (logits,) return ((loss,) + output) if loss is not None else output return CausalLMOutputWithPast(loss=loss, logits=logits)