(
self,
pretrained: str,
**kwargs,
)
| 61 | tokenizer=self.tokenizer) |
| 62 | |
| 63 | def _create_model( |
| 64 | self, |
| 65 | pretrained: str, |
| 66 | **kwargs, |
| 67 | ) -> None: |
| 68 | model_list = get_model(self.model_provider, |
| 69 | model_type=ModelType.encoder_or_decoder, |
| 70 | wrap_with_ddp=False) |
| 71 | |
| 72 | if pretrained is not None: |
| 73 | load_checkpoint(model_list, None, None) |
| 74 | |
| 75 | self._model = model_list[0] |
| 76 | |
| 77 | def tie_weights(self): |
| 78 | pass |
| 79 | self._model.tie_weights = types.MethodType(tie_weights, self._model) |
| 80 | |
| 81 | return None |
| 82 | |
| 83 | def create_model_inputs(self, tokens): |
| 84 | attention_mask, loss_mask, position_ids = get_ltor_masks_and_position_ids( |
nothing calls this directly
no test coverage detected