(self, spell_length, hidden_size, spell_func)
| 4 | |
| 5 | class PromptSpell(torch.nn.Module): |
| 6 | def __init__(self, spell_length, hidden_size, spell_func): |
| 7 | super(PromptSpell, self).__init__() |
| 8 | self.spell_length = spell_length |
| 9 | self.hidden_size = hidden_size |
| 10 | self.spell_embeddings = torch.nn.Embedding(self.spell_length, self.hidden_size) |
| 11 | self.spell_func = spell_func |
| 12 | if self.spell_func == "lstm": |
| 13 | self.lstm_head = torch.nn.LSTM(input_size=self.hidden_size, |
| 14 | hidden_size=self.hidden_size, |
| 15 | num_layers=2, |
| 16 | # dropout=self.lstm_dropout, |
| 17 | bidirectional=True, |
| 18 | batch_first=True) # .to(torch.device("cuda")) |
| 19 | self.mlp_head = torch.nn.Sequential(torch.nn.Linear(2 * self.hidden_size, self.hidden_size), |
| 20 | torch.nn.ReLU(), |
| 21 | torch.nn.Linear(self.hidden_size, self.hidden_size)) |
| 22 | elif self.spell_func == "mlp": |
| 23 | self.mlp_head = torch.nn.Sequential(torch.nn.Linear(self.hidden_size, self.hidden_size), |
| 24 | torch.nn.ReLU(), |
| 25 | torch.nn.Linear(self.hidden_size, self.hidden_size)) |
| 26 | elif self.spell_func != "none": |
| 27 | raise NotImplementedError("Prompt function " + self.spell_func) |
| 28 | |
| 29 | def init_embedding(self, word_embeddings=None, task_tokens=None): |
| 30 | num_words = 5000 |
nothing calls this directly
no outgoing calls
no test coverage detected