| 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) |
|
|