| 227 | |
| 228 | |
| 229 | def parse_args(args: list[str] | None = None): |
| 230 | parser = argparse.ArgumentParser(description="Tool for interacting with tasks") |
| 231 | parser.add_argument( |
| 232 | "TASK_FAMILY_NAME", help="The name of the task family module to import" |
| 233 | ) |
| 234 | parser.add_argument( |
| 235 | "TASK_NAME", |
| 236 | nargs="?", |
| 237 | help="The name of the task to run (required for certain operations)", |
| 238 | ) |
| 239 | parser.add_argument( |
| 240 | "OPERATION", |
| 241 | choices=[op.value for op in Operation], |
| 242 | help="The operation to perform", |
| 243 | ) |
| 244 | parser.add_argument( |
| 245 | "-s", "--submission", required=False, help="The submission string for scoring" |
| 246 | ) |
| 247 | parser.add_argument( |
| 248 | "--score_log", |
| 249 | required=False, |
| 250 | help="The JSON-encoded list of intermediate scores, or the path to a score log", |
| 251 | ) |
| 252 | parsed_args = {k.lower(): v for k, v in vars(parser.parse_args(args)).items()} |
| 253 | if ( |
| 254 | parsed_args["task_name"] is None |
| 255 | and parsed_args["operation"] not in NO_TASK_COMMANDS |
| 256 | ): |
| 257 | parser.error( |
| 258 | f"TASK_NAME is required for operation '{parsed_args['operation']}'" |
| 259 | ) |
| 260 | return parsed_args |
| 261 | |
| 262 | |
| 263 | if __name__ == "__main__": |