(self, transformer)
| 296 | |
| 297 | class OutputLayer(nn.Module): |
| 298 | def __init__(self, transformer): |
| 299 | super().__init__() |
| 300 | self.transformer = [transformer] |
| 301 | self.scale_shift_table = transformer.scale_shift_table |
| 302 | self.norm_out = transformer.norm_out |
| 303 | self.proj_out = transformer.proj_out |
| 304 | |
| 305 | @torch.autocast('cuda', dtype=AUTOCAST_DTYPE) |
| 306 | def forward(self, inputs): |