(self, encoder)
| 225 | |
| 226 | class TorchModule(Module, GenerationMixin): |
| 227 | def __init__(self, encoder): |
| 228 | super().__init__() |
| 229 | self.encoder = encoder |
| 230 | # Use hardcoded value to extend compatibility with older HF versions. |
| 231 | self.main_input_name = "input_ids" |
| 232 | |
| 233 | def forward(self, *input, **kwargs): |
| 234 | return self.encoder(*input, **kwargs)[0] |