| 182 | |
| 183 | |
| 184 | def build_model(cfg: DictConfig, num_labels: int, multiple_choice: bool = False, **kwargs): |
| 185 | if cfg.name == "hf_bert": |
| 186 | return hf_bert_module.create_hf_bert_classification( |
| 187 | num_labels=num_labels, |
| 188 | pretrained_model_name=cfg.pretrained_model_name, |
| 189 | use_pretrained=cfg.get("use_pretrained", False), |
| 190 | model_config=cfg.get("model_config", None), |
| 191 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 192 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 193 | multiple_choice=multiple_choice, |
| 194 | **kwargs, |
| 195 | ) |
| 196 | elif cfg.name == "mosaic_bert": |
| 197 | return mosaic_bert_module.create_mosaic_bert_classification( |
| 198 | num_labels=num_labels, |
| 199 | pretrained_model_name=cfg.pretrained_model_name, |
| 200 | pretrained_checkpoint=cfg.get("pretrained_checkpoint", None), |
| 201 | model_config=cfg.get("model_config", None), |
| 202 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 203 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 204 | multiple_choice=multiple_choice, |
| 205 | **kwargs, |
| 206 | ) |
| 207 | elif cfg.name == "flex_bert": |
| 208 | return flex_bert_module.create_flex_bert_classification( |
| 209 | num_labels=num_labels, |
| 210 | pretrained_model_name=cfg.pretrained_model_name, |
| 211 | pretrained_checkpoint=cfg.get("pretrained_checkpoint", None), |
| 212 | model_config=cfg.get("model_config", None), |
| 213 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 214 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 215 | multiple_choice=multiple_choice, |
| 216 | **kwargs, |
| 217 | ) |
| 218 | else: |
| 219 | raise ValueError(f"Not sure how to build model with name={cfg.name}") |
| 220 | |
| 221 | |
| 222 | def get_values_from_path(path: str, separator: str = "/") -> Dict[str, str]: |