(cfg: DictConfig)
| 302 | |
| 303 | |
| 304 | def build_model(cfg: DictConfig): |
| 305 | if cfg.name == "hf_bert": |
| 306 | return hf_bert_module.create_hf_bert_mlm( |
| 307 | pretrained_model_name=cfg.pretrained_model_name, |
| 308 | use_pretrained=cfg.get("use_pretrained", None), |
| 309 | model_config=cfg.get("model_config", None), |
| 310 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 311 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 312 | ) |
| 313 | elif cfg.name == "mosaic_bert": |
| 314 | return mosaic_bert_module.create_mosaic_bert_mlm( |
| 315 | pretrained_model_name=cfg.pretrained_model_name, |
| 316 | pretrained_checkpoint=cfg.get("pretrained_checkpoint", None), |
| 317 | model_config=cfg.get("model_config", None), |
| 318 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 319 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 320 | ) |
| 321 | elif cfg.name == "flex_bert": |
| 322 | return flex_bert_module.create_flex_bert_mlm( |
| 323 | pretrained_model_name=cfg.pretrained_model_name, |
| 324 | pretrained_checkpoint=cfg.get("pretrained_checkpoint", None), |
| 325 | model_config=cfg.get("model_config", None), |
| 326 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 327 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 328 | recompute_metric_loss=cfg.get("recompute_metric_loss", False), |
| 329 | disable_train_metrics=cfg.get("disable_train_metrics", False), |
| 330 | ) |
| 331 | else: |
| 332 | raise ValueError(f"Not sure how to build model with name={cfg.name}") |
| 333 | |
| 334 | |
| 335 | def init_from_checkpoint(cfg: DictConfig, new_model: nn.Module): |
no outgoing calls
no test coverage detected