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

Function main

generate_eval_config.py:209–427  ·  view source on GitHub ↗
(
    checkpoint: Annotated[Path, Option(help="Path to a model checkpoint", show_default=False, rich_help_panel="Checkpoint & Config Paths")],
    output_dir: Annotated[Path, Option(help="Output directory for the generated config", rich_help_panel="Checkpoint & Config Paths")] = Path("./yamls/ablations"),
    train_config: Annotated[Optional[Path], Option(help="Path to a .yaml file containing training configuration. If one is not provided, will attempt to load the config from a wandb run or use defaults.", rich_help_panel="Checkpoint & Config Paths")] = None,
    model_size: Annotated[ModelSize, Option("--model-size", help="Model to use for default model config: 'base' or 'large'", rich_help_panel="Checkpoint & Config Paths")] = ModelSize.BASE,
    rope_theta: Annotated[Optional[float], Option("--rope-theta", help="Value for `rotary_emb_base` in the model configuration. If not provided, defaults to pretraining value of 10000.0", rich_help_panel="Checkpoint & Config Paths")] = None,
    use_dir_name: Annotated[bool, Option("--use-dir-name", help="Use the checkpoint's parent dirname as the eval base_run_name", rich_help_panel="Checkpoint & Config Paths")] = False,
    tasks: Annotated[Optional[List[TaskName]], Option(help="List of tasks to include in the evaluation. Default is all tasks.", rich_help_panel="Eval Tasks", case_sensitive=False, show_default=False)] = None, # type: ignore
    wandb_run: Annotated[Optional[str], Option(help="wandb run containing the training configuration", rich_help_panel="Weights & Biases")] = None,
    wandb_project: Annotated[Optional[str], Option(help="wandb project for the run", rich_help_panel="Weights & Biases")] = None,
    wandb_entity: Annotated[Optional[str], Option(help="wandb entity for the project", rich_help_panel="Weights & Biases")] = None,
    track_run: Annotated[bool, Option("--track-run", help="Track the eval run with wandb", rich_help_panel="Weights & Biases")] = False,
    track_run_project: Annotated[Optional[str], Option(help="wandb project for tracking the run", rich_help_panel="Weights & Biases")] = None,
    pooling_type: Annotated[Optional[str], Option(help="Pooling type for the classification head", show_default=False, rich_help_panel="Model Options")] = None,
    head_class_act: Annotated[Optional[str], Option(help="Classification head activation function. Defaults to hidden_act if set, then tanh", show_default=False, rich_help_panel="Model Options")] = None,
    head_class_norm: Annotated[Optional[str], Option(help="Classification head normalization function", show_default=False, rich_help_panel="Model Options")] = None,
    head_class_dropout: Annotated[float, Option(help="Classification head dropout rate", rich_help_panel="Model Options")] = 0.0,
    fast_ultrafeedback: Annotated[bool, Option("--fast-ultrafeedback", help="Use a shorter sequence length (1536) for the UltraFeedback eval", rich_help_panel="Task Settings")] = False,
    seeds: Annotated[List[int], Option(help="List of seeds to use for the eval", rich_help_panel="Task Settings")] = [1618, 42, 6033, 3145],
    parallel: Annotated[bool, Option("--parallel/--single", help="Run the evals in parallel on multiple GPUs or one GPU. Only use if evaluating a single checkpoint on multiple GPUs.", rich_help_panel="Task Settings")] = False,
    config: Annotated[Optional[Path], Option(callback=conf_callback, is_eager=True, help="Relative path to YAML config file for setting options. Passing CLI options will supersede config options.", case_sensitive=False, rich_help_panel="Options")] = None,
)

Source from the content-addressed store, hash-verified

207
208@app.command()
209def main(
210 checkpoint: Annotated[Path, Option(help="Path to a model checkpoint", show_default=False, rich_help_panel="Checkpoint & Config Paths")],
211 output_dir: Annotated[Path, Option(help="Output directory for the generated config", rich_help_panel="Checkpoint & Config Paths")] = Path("./yamls/ablations"),
212 train_config: Annotated[Optional[Path], Option(help="Path to a .yaml file containing training configuration. If one is not provided, will attempt to load the config from a wandb run or use defaults.", rich_help_panel="Checkpoint & Config Paths")] = None,
213 model_size: Annotated[ModelSize, Option("--model-size", help="Model to use for default model config: 'base' or 'large'", rich_help_panel="Checkpoint & Config Paths")] = ModelSize.BASE,
214 rope_theta: Annotated[Optional[float], Option("--rope-theta", help="Value for `rotary_emb_base` in the model configuration. If not provided, defaults to pretraining value of 10000.0", rich_help_panel="Checkpoint & Config Paths")] = None,
215 use_dir_name: Annotated[bool, Option("--use-dir-name", help="Use the checkpoint's parent dirname as the eval base_run_name", rich_help_panel="Checkpoint & Config Paths")] = False,
216 tasks: Annotated[Optional[List[TaskName]], Option(help="List of tasks to include in the evaluation. Default is all tasks.", rich_help_panel="Eval Tasks", case_sensitive=False, show_default=False)] = None, # type: ignore
217 wandb_run: Annotated[Optional[str], Option(help="wandb run containing the training configuration", rich_help_panel="Weights & Biases")] = None,
218 wandb_project: Annotated[Optional[str], Option(help="wandb project for the run", rich_help_panel="Weights & Biases")] = None,
219 wandb_entity: Annotated[Optional[str], Option(help="wandb entity for the project", rich_help_panel="Weights & Biases")] = None,
220 track_run: Annotated[bool, Option("--track-run", help="Track the eval run with wandb", rich_help_panel="Weights & Biases")] = False,
221 track_run_project: Annotated[Optional[str], Option(help="wandb project for tracking the run", rich_help_panel="Weights & Biases")] = None,
222 pooling_type: Annotated[Optional[str], Option(help="Pooling type for the classification head", show_default=False, rich_help_panel="Model Options")] = None,
223 head_class_act: Annotated[Optional[str], Option(help="Classification head activation function. Defaults to hidden_act if set, then tanh", show_default=False, rich_help_panel="Model Options")] = None,
224 head_class_norm: Annotated[Optional[str], Option(help="Classification head normalization function", show_default=False, rich_help_panel="Model Options")] = None,
225 head_class_dropout: Annotated[float, Option(help="Classification head dropout rate", rich_help_panel="Model Options")] = 0.0,
226 fast_ultrafeedback: Annotated[bool, Option("--fast-ultrafeedback", help="Use a shorter sequence length (1536) for the UltraFeedback eval", rich_help_panel="Task Settings")] = False,
227 seeds: Annotated[List[int], Option(help="List of seeds to use for the eval", rich_help_panel="Task Settings")] = [1618, 42, 6033, 3145],
228 parallel: Annotated[bool, Option("--parallel/--single", help="Run the evals in parallel on multiple GPUs or one GPU. Only use if evaluating a single checkpoint on multiple GPUs.", rich_help_panel="Task Settings")] = False,
229 config: Annotated[Optional[Path], Option(callback=conf_callback, is_eager=True, help="Relative path to YAML config file for setting options. Passing CLI options will supersede config options.", case_sensitive=False, rich_help_panel="Options")] = None,
230): # fmt: skip
231 # Read the input YAML file
232 output_dir.mkdir(parents=True, exist_ok=True)
233 input_config = None
234
235 if checkpoint.is_file() and checkpoint.suffix == ".pt":
236 ckpt = checkpoint.name # checkpoint
237 ckpt_path = checkpoint.parent
238 elif checkpoint.is_dir():
239 ckpts = list(checkpoint.glob("*.pt"))
240 if len(ckpts) == 1:
241 ckpt = ckpts[0].name
242 elif len(ckpts) > 1:
243 ckpt = "latest-rank0.pt"
244 elif len(ckpts) == 0:
245 raise ValueError(f"No checkpoint found in the provided directory: {checkpoint}")
246 ckpt_path = checkpoint
247 else:
248 raise ValueError(f"Invalid checkpoint path provided: {checkpoint}")
249
250 ckpt_id = ckpt_path.name
251
252 if train_config:
253 with train_config.open("r") as file:
254 input_config = yaml.safe_load(file)
255 else:
256 # Specify the run name
257 print("Attempting to find config file within checkpoint folder...")
258 yaml_file = checkpoint.parent / f"{checkpoint.parent.name}.yaml"
259 yaml_file_alt = ckpt_path / f"{ckpt_id}.yaml"
260 print(yaml_file)
261
262 if yaml_file.exists():
263 with yaml_file.open("r") as file:
264 input_config = yaml.safe_load(file)
265 elif yaml_file_alt.exists():
266 with yaml_file_alt.open("r") as file:

Callers

nothing calls this directly

Calls 4

get_wandb_configFunction · 0.85
safe_getFunction · 0.85
get_model_defaultsFunction · 0.85
ordered_yaml_dumpFunction · 0.85

Tested by

no test coverage detected