OddEvenDumb / modeling_oddevendumb.py
56m's picture
Upload 7 files
e5a9af5 verified
Raw
History Blame Contribute Delete
2.17 kB
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)