(arg_name, old_arg_name=None)
| 46 | args = get_args() |
| 47 | |
| 48 | def _compare(arg_name, old_arg_name=None): |
| 49 | if old_arg_name is not None: |
| 50 | checkpoint_value = getattr(checkpoint_args, old_arg_name) |
| 51 | else: |
| 52 | checkpoint_value = getattr(checkpoint_args, arg_name) |
| 53 | args_value = getattr(args, arg_name) |
| 54 | error_message = ( |
| 55 | "{} value from checkpoint ({}) is not equal to the " |
| 56 | "input argument value ({}).".format(arg_name, checkpoint_value, args_value) |
| 57 | ) |
| 58 | assert checkpoint_value == args_value, error_message |
| 59 | |
| 60 | _compare("num_layers") |
| 61 | _compare("hidden_size") |
no outgoing calls
no test coverage detected