r""" The ValueHead class implements a head for GPT2 that returns a scalar for each output token.
| 26 | |
| 27 | |
| 28 | class ValueHead(nn.Module): |
| 29 | r""" |
| 30 | The ValueHead class implements a head for GPT2 that returns a scalar for each output token. |
| 31 | """ |
| 32 | |
| 33 | def __init__(self, config, **kwargs): |
| 34 | super().__init__() |
| 35 | if not hasattr(config, "summary_dropout_prob"): |
| 36 | summary_dropout_prob = kwargs.pop("summary_dropout_prob", 0.1) |
| 37 | else: |
| 38 | summary_dropout_prob = config.summary_dropout_prob |
| 39 | |
| 40 | self.dropout = nn.Dropout(summary_dropout_prob) if summary_dropout_prob else nn.Identity() |
| 41 | |
| 42 | # some models such as OPT have a projection layer before the word embeddings - e.g. OPT-350m |
| 43 | if hasattr(config, "hidden_size"): |
| 44 | hidden_size = config.hidden_size |
| 45 | if hasattr(config, "word_embed_proj_dim"): |
| 46 | hidden_size = config.word_embed_proj_dim |
| 47 | elif hasattr(config, "is_encoder_decoder"): |
| 48 | if config.is_encoder_decoder and hasattr(config, "decoder"): |
| 49 | if hasattr(config.decoder, "hidden_size"): |
| 50 | hidden_size = config.decoder.hidden_size |
| 51 | |
| 52 | self.summary = nn.Linear(hidden_size, 1) |
| 53 | |
| 54 | self.flatten = nn.Flatten() |
| 55 | |
| 56 | def forward(self, hidden_states): |
| 57 | |
| 58 | if hidden_states.device != self.summary.weight.device: |
| 59 | hidden_states = hidden_states.to(self.summary.weight.device) |
| 60 | |
| 61 | output = self.dropout(hidden_states) |
| 62 | |
| 63 | # For now force upcast in fp32 if needed. Let's keep the |
| 64 | # output in fp32 for numerical stability. |
| 65 | if output.dtype != self.summary.weight.dtype: |
| 66 | output = output.to(self.summary.weight.dtype) |
| 67 | |
| 68 | output = self.summary(output) |
| 69 | values = torch.tanh(output).squeeze(-1) |
| 70 | return values |
| 71 | |
| 72 | |
| 73 | class AutoModelForCausalLMWithValueHead(nn.Module): |
no outgoing calls
no test coverage detected