37 use num_types,
only: rp
38 use case,
only: case_t
39 use json_file_module,
only: json_file
40 use json_utils,
only: json_get, json_get_or_default
41 use chkp_output,
only: chkp_output_t
42 use field,
only: field_t
43 use field_list,
only: field_list_t
44 use logger,
only: neko_log, log_size, neko_log_debug
45 use utils,
only: neko_error, mkdir
46 use math,
only: copy, rzero
47 use host_array,
only: host_array_t
48 use profiler,
only: profiler_start_region, profiler_end_region
50 use comm,
only: pe_rank, neko_comm
51 use neko_config,
only: neko_bcknd_device
52 use device,
only: device_memcpy, device_to_host, host_to_device
53 use registry,
only: neko_registry
54 use mpi_f08,
only: mpi_barrier
59 integer,
parameter :: CHECKPOINT_LINEAR = 1
68 character(len=256) :: algorithm =
""
70 character(len=256) :: filename =
"forward_checkpoint"
72 character(len=256) :: path =
"checkpoints/"
74 character(len=8) :: fmt =
"chkp"
76 integer :: n_saves_memory = 10
78 logical :: keep_checkpoints = .false.
81 integer :: algorithm_id = 0
82 integer :: n_saves_disc = 0
83 integer :: first_valid_timestep = 2
84 integer :: loaded_checkpoint = -1
87 type(field_list_t) :: state_list
88 type(host_array_t),
dimension(:,:),
allocatable :: state_storage
91 type(chkp_output_t) :: chkp_output
97 procedure,
public, pass(this) :: init_from_components => &
98 checkpoint_init_from_components
100 procedure,
public, pass(this) :: free => checkpoint_free
102 procedure,
public, pass(this) :: reset => checkpoint_reset
104 procedure,
public, pass(this) ::
save => checkpoint_save
106 procedure,
public, pass(this) :: restore => checkpoint_restore
109 procedure, pass(this) :: save_data => checkpoint_save_data
111 procedure, pass(this) :: load_data => checkpoint_load_data
119 module subroutine checkpoint_save_linear(this)
120 class(state_recover_checkpoint_t),
intent(inout) :: this
121 end subroutine checkpoint_save_linear
124 module subroutine checkpoint_restore_linear(this, tstep)
125 class(state_recover_checkpoint_t),
intent(inout) :: this
126 integer,
intent(in) :: tstep
127 end subroutine checkpoint_restore_linear
140 subroutine checkpoint_init_from_json(this, neko_case, params)
141 class(state_recover_checkpoint_t),
intent(inout) :: this
142 class(case_t),
target,
intent(inout) :: neko_case
143 type(json_file),
intent(inout) :: params
144 integer :: n_saves_memory
145 character(len=:),
allocatable :: path, filename, algorithm, fmt
146 character(len=256),
dimension(:),
allocatable :: extra_field_names
147 type(field_list_t) :: extra_fields
148 type(field_t),
pointer :: fi
150 logical :: keep_checkpoints
152 call json_get_or_default(params,
"algorithm", algorithm,
"linear")
153 call json_get_or_default(params,
"n_memory", n_saves_memory, 10)
154 call json_get_or_default(params,
"path", path,
"checkpoints/")
155 call json_get_or_default(params,
"filename", filename,
"checkpoint")
156 call json_get_or_default(params,
"format", fmt,
"chkp")
157 call json_get_or_default(params,
"keep_checkpoints", keep_checkpoints, &
160 if (
"extra_fields" .in. params)
then
161 call json_get(params,
"extra_fields", extra_field_names)
162 call extra_fields%init(
size(extra_field_names))
163 do i = 1,
size(extra_field_names)
164 fi => neko_registry%get_field(extra_field_names(i))
165 call extra_fields%assign(i, fi)
168 call this%init_from_components(neko_case, algorithm, n_saves_memory, &
169 path, filename, fmt, keep_checkpoints, extra_fields)
172 call this%init_from_components(neko_case, algorithm, n_saves_memory, &
173 path, filename, fmt, keep_checkpoints)
176 end subroutine checkpoint_init_from_json
188 subroutine checkpoint_init_from_components(this, neko_case, algorithm, &
189 n_saves_memory, path, filename, fmt, keep_checkpoints, extra_fields)
190 class(state_recover_checkpoint_t),
intent(inout),
target :: this
191 class(case_t),
target,
intent(inout) :: neko_case
192 character(len=*),
intent(in) :: algorithm
193 integer,
optional,
intent(in) :: n_saves_memory
194 character(len=*),
optional,
intent(in) :: path
195 character(len=*),
optional,
intent(in) :: filename
196 character(len=*),
optional,
intent(in) :: fmt
197 logical,
optional,
intent(in) :: keep_checkpoints
198 type(field_list_t),
optional,
intent(inout) :: extra_fields
199 type(field_t),
pointer :: si
200 character(len=LOG_SIZE) :: msg
201 integer :: i, j, n_states
205 this%neko_case => neko_case
208 if (
present(n_saves_memory)) this%n_saves_memory = n_saves_memory
209 if (
present(path)) this%path = trim(path)
210 if (
present(filename)) this%filename = trim(filename)
211 if (
present(fmt)) this%fmt = trim(fmt)
212 if (
present(keep_checkpoints)) this%keep_checkpoints = keep_checkpoints
215 select case (trim(algorithm))
216 case (
"linear",
"LINEAR",
"Linear")
217 this%algorithm =
"linear"
218 this%algorithm_id = checkpoint_linear
220 call neko_error(
"Only the linear checkpoint strategy is supported.")
223 inquire(file = trim(this%path), exist = exists)
224 if (.not. exists)
then
225 call mpi_barrier(neko_comm)
226 if (pe_rank .eq. 0)
then
227 call mkdir(trim(this%path))
229 call mpi_barrier(neko_comm)
233 call this%chkp_output%init(neko_case%chkp, this%filename, &
234 fmt = this%fmt, path = this%path, overwrite = .true.)
237 if (
allocated(neko_case%scalars))
then
238 n_states = n_states +
size(neko_case%scalars%scalar_fields)
240 if (
present(extra_fields))
then
241 n_states = n_states + extra_fields%size()
244 call this%state_list%init(n_states)
247 call this%state_list%assign(1, neko_case%fluid%p)
248 call this%state_list%assign(2, neko_case%fluid%u)
249 call this%state_list%assign(3, neko_case%fluid%v)
250 call this%state_list%assign(4, neko_case%fluid%w)
254 if (
allocated(neko_case%scalars))
then
255 do i = 1,
size(neko_case%scalars%scalar_fields)
256 si => neko_case%scalars%scalar_fields(i)%scalar%s
257 call this%state_list%assign(n_states + i, si)
259 n_states = n_states +
size(neko_case%scalars%scalar_fields)
263 if (
present(extra_fields))
then
264 do i = 1, extra_fields%size()
265 si => extra_fields%get_by_index(i)
266 call this%state_list%assign(n_states + i, si)
268 n_states = n_states + extra_fields%size()
272 allocate(this%state_storage(this%n_saves_memory, this%state_list%size()))
273 do i = 1, this%n_saves_memory
274 do j = 1, this%state_list%size()
275 si => this%state_list%get(j)
276 call this%state_storage(i, j)%init(si%size())
281 call neko_log%section(
"Checkpointing")
283 write(msg,
'(A, A)')
"Algorithm: ", trim(this%algorithm)
284 call neko_log%message(trim(msg))
285 write(msg,
'(A,I0)')
"Number of checkpoints in RAM: ", this%n_saves_memory
286 call neko_log%message(trim(msg))
287 write(msg,
'(A, A)')
"Checkpoint file path: ", trim(this%path)
288 call neko_log%message(trim(msg))
289 write(msg,
'(A, A)')
"Checkpoint file name: ", trim(this%filename)
290 call neko_log%message(trim(msg))
291 write(msg,
'(A, A)')
"Checkpoint file format: ", trim(this%fmt)
292 call neko_log%message(trim(msg))
294 if (.not. this%keep_checkpoints)
then
295 call neko_log%message(
"Checkpoint files will be deleted.")
297 call neko_log%message(
"Checkpoint files will be kept.")
300 call neko_log%message(
"Fields in checkpoint:", neko_log_debug)
301 do i = 1, this%state_list%size()
302 si => this%state_list%get(i)
303 call neko_log%message(
" - " // trim(si%name), neko_log_debug)
306 call neko_log%end_section()
308 end subroutine checkpoint_init_from_components
312 subroutine checkpoint_free(this)
313 class(state_recover_checkpoint_t),
intent(inout) :: this
315 character(len=1024) :: file_name
317 integer :: stat, unit
320 if (
allocated(this%state_storage))
then
321 do i = 1, this%n_saves_memory
322 do j = 1, this%state_list%size()
323 call this%state_storage(i, j)%free()
328 call this%state_list%free()
329 if (
allocated(this%state_storage))
deallocate(this%state_storage)
332 if (.not. this%keep_checkpoints .and. pe_rank .eq. 0)
then
333 do i = this%get_n_timesteps(), 1, -1
334 call this%chkp_output%set_counter(i)
335 file_name = this%chkp_output%file_%get_fname()
336 inquire(file = trim(file_name), exist = exists)
338 open(newunit = unit, file = trim(file_name), iostat = stat, &
340 if (stat .eq. 0)
close(unit, status =
'delete')
344 call mpi_barrier(neko_comm)
347 this%filename =
"checkpoint"
349 this%algorithm =
"linear"
350 this%n_saves_memory = 10
351 this%keep_checkpoints = .false.
353 this%n_saves_disc = 0
354 call this%set_n_timesteps(0)
355 this%first_valid_timestep = 2
356 this%loaded_checkpoint = -1
357 nullify(this%neko_case)
359 end subroutine checkpoint_free
367 subroutine checkpoint_save(this)
368 class(state_recover_checkpoint_t),
intent(inout) :: this
370 call profiler_start_region(
"Checkpoint save")
373 call this%set_n_timesteps(this%get_n_timesteps() + 1)
375 select case (this%algorithm_id)
376 case (checkpoint_linear)
377 call checkpoint_save_linear(this)
379 call neko_error(
"Unknown checkpoint algorithm: " // this%algorithm)
382 call profiler_end_region(
"Checkpoint save")
383 end subroutine checkpoint_save
389 subroutine checkpoint_restore(this, tstep)
390 class(state_recover_checkpoint_t),
intent(inout) :: this
391 integer,
intent(in) :: tstep
392 character(len=256) :: msg
394 call profiler_start_region(
"Checkpoint restore")
396 if (tstep .lt. 1 .or. tstep .gt. this%get_n_timesteps())
then
397 write(msg,
'(A,I0,A,I0,A)')
"Requested timestep ", tstep, &
398 " is out of range [1, ", this%get_n_timesteps(),
"]"
399 call neko_error(trim(msg))
402 select case (this%algorithm_id)
403 case (checkpoint_linear)
404 call checkpoint_restore_linear(this, tstep)
406 call neko_error(
"Unknown checkpoint algorithm: " // this%algorithm)
409 call profiler_end_region(
"Checkpoint restore")
410 end subroutine checkpoint_restore
415 subroutine checkpoint_save_data(this, index)
416 class(state_recover_checkpoint_t),
intent(inout) :: this
417 integer,
intent(in) :: index
418 type(field_t),
pointer :: si
420 character(len=1024) :: msg
422 if (index .lt. 1 .or. index .gt. this%n_saves_memory)
then
423 write(msg,
'(A,I0,A,I0,A)')
"Checkpoint save index ", index, &
424 " is out of range [1, ", this%n_saves_memory,
"]"
425 call neko_error(trim(msg))
429 do i = 1, this%state_list%size()
430 if (.not. this%state_storage(index, i)%is_allocated())
then
431 si => this%state_list%get(i)
432 call this%state_storage(index, i)%init(si%size())
437 if (neko_bcknd_device .eq. 0)
then
438 do i = 1, this%state_list%size()
439 si => this%state_list%get(i)
440 call copy(this%state_storage(index, i)%x, si%x, si%size())
443 do i = 1, this%state_list%size()
444 si => this%state_list%get(i)
445 call device_memcpy(this%state_storage(index, i)%x, si%x_d, &
446 si%size(), device_to_host, this%state_list%size() .eq. i)
451 end subroutine checkpoint_save_data
457 subroutine checkpoint_load_data(this, index)
458 class(state_recover_checkpoint_t),
intent(inout) :: this
459 integer,
intent(in) :: index
460 type(field_t),
pointer :: si
461 character(len=1024) :: msg
464 if (index .lt. 1 .or. index .gt. this%n_saves_memory)
then
465 write(msg,
'(A,I0,A,I0,A)')
"Checkpoint save index ", index, &
466 " is out of range [1, ", this%n_saves_memory,
"]"
467 call neko_error(trim(msg))
471 if (neko_bcknd_device .eq. 0)
then
472 do i = 1, this%state_list%size()
473 si => this%state_list%get(i)
474 call copy(si%x, this%state_storage(index, i)%x, si%size())
477 do i = 1, this%state_list%size()
478 si => this%state_list%get(i)
479 call device_memcpy(this%state_storage(index, i)%x, si%x_d, &
480 si%size(), host_to_device, this%state_list%size() .eq. i)
485 end subroutine checkpoint_load_data
492 subroutine checkpoint_reset(this)
493 class(state_recover_checkpoint_t),
intent(inout) :: this
497 this%loaded_checkpoint = -1
498 this%n_saves_disc = 0
499 call this%set_n_timesteps(0)
501 do i = 1,
size(this%state_storage, 1)
502 do j = 1,
size(this%state_storage, 2)
503 call rzero(this%state_storage(i, j)%x, this%state_storage(i, j)%size())
507 end subroutine checkpoint_reset
Checkpoint-based state recovery for adjoint runs.
subroutine checkpoint_init_from_json(this, neko_case, params)
Save the current state of the simulation in a linear fashion.
Abstract interface for state recovery strategies.
Abstract base type for state recovery implementations.