MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / build_model

Function build_model

main.py:304–332  ·  view source on GitHub ↗
(cfg: DictConfig)

Source from the content-addressed store, hash-verified

302
303
304def 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
335def init_from_checkpoint(cfg: DictConfig, new_model: nn.Module):

Callers 2

init_from_checkpointFunction · 0.70
mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected