Repository navigation
Expand file tree
/
Copy pathrun.py
More file actions
63 lines (52 loc) · 2.31 KB
/
Copy pathrun.py
File metadata and controls
63 lines (52 loc) · 2.31 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
import argparse
import os
os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false"
import yaml
# Parse args early to determine solver type before importing JAX-heavy modules
parser = argparse.ArgumentParser(description="Automatic Differentiation Enabled Plasma Transport")
parser.add_argument("--cfg", help="enter path to cfg")
parser.add_argument("--run_id", help="enter run_id to continue")
parser.add_argument("--dt", help="override grid.dt (for stability/timing runs)")
parser.add_argument("--tmax", help="override grid.tmax and save.t.tmax (for smoke/timing runs)")
parser.add_argument("--save-nt", type=int, help="override save.t.nt")
parser.add_argument("--run-name", help="override mlflow.run")
parser.add_argument("--disable-ib", action="store_true", help="disable inverse-bremsstrahlung heating")
parser.add_argument(
"--disable-hidden-density-gradient",
action="store_true",
help="disable the kinetic-Ohm hidden density gradient",
)
args = parser.parse_args()
# Enable float64 for kinetic solvers (must be done before importing adept)
if args.run_id is None and args.cfg:
with open(f"{os.path.join(os.getcwd(), args.cfg)}.yaml") as fi:
cfg = yaml.safe_load(fi)
if args.dt is not None:
cfg["grid"]["dt"] = args.dt
if args.tmax is not None:
cfg["grid"]["tmax"] = args.tmax
cfg.setdefault("save", {}).setdefault("t", {})["tmax"] = args.tmax
if args.save_nt is not None:
cfg.setdefault("save", {}).setdefault("t", {})["nt"] = args.save_nt
if args.run_name is not None:
cfg.setdefault("mlflow", {})["run"] = args.run_name
if args.disable_ib:
cfg.get("drivers", {}).pop("ib", None)
if args.disable_hidden_density_gradient:
hidden = cfg.get("terms", {}).get("field_solver", {}).get("hidden_density_gradient")
if hidden is not None:
hidden["active"] = False
if cfg.get("solver") != "envelope-2d":
from jax import config
config.update("jax_enable_x64", True)
# Now safe to import adept (which imports JAX)
from adept import ergoExo
if __name__ == "__main__":
exo = ergoExo()
if args.run_id is None:
# Config already loaded above
modules = exo.setup(cfg=cfg)
sol, post_out, run_id = exo(modules)
else:
exo.run_job(args.run_id, nested=None)
run_id = args.run_id