(checkpoints_path,
iteration,
release=False,
zero=False)
| 178 | |
| 179 | |
| 180 | def get_checkpoint_name(checkpoints_path, |
| 181 | iteration, |
| 182 | release=False, |
| 183 | zero=False): |
| 184 | if release: |
| 185 | d = 'release' |
| 186 | else: |
| 187 | d = '{}'.format(iteration) |
| 188 | if zero: |
| 189 | dp_rank = mpu.get_data_parallel_rank() |
| 190 | d += '_zero_dp_rank_{}'.format(dp_rank) |
| 191 | return os.path.join( |
| 192 | checkpoints_path, d, |
| 193 | 'mp_rank_{:02d}_model_states.pt'.format(mpu.get_model_parallel_rank())) |
| 194 | |
| 195 | |
| 196 | def ensure_directory_exists(filename): |
no outgoing calls
no test coverage detected
searching dependent graphs…