mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-08-05 11:07:43 +00:00
fix: move snapshot saving after step index update in loop execution (#1206)
This commit is contained in:
@@ -221,6 +221,7 @@ class LoopBase:
|
||||
# NOTE: each step are aware are of current loop index
|
||||
# It is very important to set it before calling the step function!
|
||||
self.loop_prev_out[li][self.LOOP_IDX_KEY] = li
|
||||
|
||||
try:
|
||||
# Call function with current loop's output, await if coroutine or use ProcessPoolExecutor for sync if required
|
||||
if force_subproc:
|
||||
@@ -236,9 +237,6 @@ class LoopBase:
|
||||
result = func(self.loop_prev_out[li])
|
||||
# Store result in the nested dictionary
|
||||
self.loop_prev_out[li][name] = result
|
||||
|
||||
# Save snapshot after completing the step
|
||||
self.dump(self.session_folder / f"{li}" / f"{si}_{name}")
|
||||
except Exception as e:
|
||||
if isinstance(e, self.skip_loop_error):
|
||||
logger.warning(f"Skip loop {li} due to {e}")
|
||||
@@ -256,6 +254,8 @@ class LoopBase:
|
||||
else:
|
||||
raise # re-raise unhandled exceptions
|
||||
finally:
|
||||
# No matter the execution succeed or not, we have to finish the following steps
|
||||
|
||||
# Record the trace
|
||||
end = datetime.now(timezone.utc)
|
||||
self.loop_trace[li].append(LoopTrace(start, end, step_idx=si))
|
||||
@@ -279,6 +279,12 @@ class LoopBase:
|
||||
step_index=next_step,
|
||||
step_name=self.steps[next_step],
|
||||
)
|
||||
|
||||
# Save snapshot after completing the step;
|
||||
# 1) It has to be after the step_idx is updated, so loading the snapshot will be on the right step.
|
||||
# 2) Only save it when the step forward, withdraw does not worth saving.
|
||||
self.dump(self.session_folder / f"{li}" / f"{si}_{name}")
|
||||
|
||||
self._check_exit_conditions_on_step(loop_id=li, step_id=si)
|
||||
else:
|
||||
logger.warning(f"Step forward {si} of loop {li} is skipped.")
|
||||
|
||||
Reference in New Issue
Block a user