| 124 | |
| 125 | |
| 126 | def build_model( |
| 127 | cfg: DictConfig, num_labels: int, multiple_choice: bool = False, **kwargs |
| 128 | ): |
| 129 | if cfg.name == "hf_bert": |
| 130 | return hf_bert_module.create_hf_bert_classification( |
| 131 | num_labels=num_labels, |
| 132 | pretrained_model_name=cfg.pretrained_model_name, |
| 133 | use_pretrained=cfg.get("use_pretrained", False), |
| 134 | model_config=cfg.get("model_config", None), |
| 135 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 136 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 137 | multiple_choice=multiple_choice, |
| 138 | **kwargs, |
| 139 | ) |
| 140 | elif cfg.name == "mosaic_bert": |
| 141 | return mosaic_bert_module.create_mosaic_bert_classification( |
| 142 | num_labels=num_labels, |
| 143 | pretrained_model_name=cfg.pretrained_model_name, |
| 144 | pretrained_checkpoint=cfg.get("pretrained_checkpoint", None), |
| 145 | model_config=cfg.get("model_config", None), |
| 146 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 147 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 148 | multiple_choice=multiple_choice, |
| 149 | **kwargs, |
| 150 | ) |
| 151 | elif cfg.name == "flex_bert": |
| 152 | return flex_bert_module.create_flex_bert_classification( |
| 153 | num_labels=num_labels, |
| 154 | pretrained_model_name=cfg.pretrained_model_name, |
| 155 | pretrained_checkpoint=cfg.get("pretrained_checkpoint", None), |
| 156 | model_config=cfg.get("model_config", None), |
| 157 | tokenizer_name=cfg.get("tokenizer_name", None), |
| 158 | gradient_checkpointing=cfg.get("gradient_checkpointing", None), |
| 159 | multiple_choice=multiple_choice, |
| 160 | **kwargs, |
| 161 | ) |
| 162 | else: |
| 163 | raise ValueError(f"Not sure how to build model with name={cfg.name}") |
| 164 | |
| 165 | |
| 166 | def get_values_from_path(path: str, separator: str = "/") -> Dict[str, str]: |