diff --git a/pyro/mesh/tests/test_io.py b/pyro/mesh/tests/test_io.py index 1c8eba265..e26c980cb 100644 --- a/pyro/mesh/tests/test_io.py +++ b/pyro/mesh/tests/test_io.py @@ -3,6 +3,7 @@ import pyro.mesh.boundary as bnd import pyro.util.io_pyro as io +from pyro import Pyro from pyro.mesh import patch @@ -28,3 +29,17 @@ def test_write_read(): anew = nd.get_var("a") assert_array_equal(anew.v(), a.v()) + + +def test_simulation_step_count_round_trip(tmp_path, monkeypatch): + """Variable-name iteration must not overwrite the saved step counter.""" + monkeypatch.chdir(tmp_path) + p = Pyro("advection") + p.initialize_problem("test", inputs_dict={"mesh.nx": 8, "mesh.ny": 8, + "driver.max_steps": 3, "io.force_final_output": 0}) + p.run_sim() + p.sim.write("checkpoint") + restored = io.read("checkpoint") + assert restored.n == 3 + assert restored.cc_data.t == p.sim.cc_data.t + assert_array_equal(restored.cc_data.get_var("density").v(), p.sim.cc_data.get_var("density").v()) diff --git a/pyro/util/io_pyro.py b/pyro/util/io_pyro.py index 37f263cec..f2901c350 100644 --- a/pyro/util/io_pyro.py +++ b/pyro/util/io_pyro.py @@ -40,7 +40,7 @@ def read(filename): solver_name = f.attrs["solver"] problem_name = f.attrs["problem"] t = f.attrs["time"] - n = f.attrs["nsteps"] + nsteps = f.attrs["nsteps"] except KeyError: # this was just a patch written out solver_name = None @@ -120,7 +120,7 @@ def read(filename): solver = importlib.import_module(f"pyro.{solver_name}") sim = solver.Simulation(solver_name, problem_name, None, None) - sim.n = n + sim.n = nsteps sim.cc_data = myd sim.cc_data.t = t sim.particles = my_particles