-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsetup_imports.py
More file actions
35 lines (28 loc) · 1.05 KB
/
Copy pathsetup_imports.py
File metadata and controls
35 lines (28 loc) · 1.05 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
import jax
import jax.numpy as jnp
from jax.experimental.ode import odeint
import matplotlib.pyplot as plt
import matplotlib.animation as animation
from IPython.display import HTML
import yaml
print("JAX version:", jax.__version__)
print("JAX backend:", jax.default_backend())
key = jax.random.PRNGKey(0)
jax.config.update("jax_enable_x64", True)
print("JAX 64-bit precision enabled:", jax.config.read("jax_enable_x64"))
# Load config.yaml
with open("config.yaml", "r") as f:
config = yaml.safe_load(f)
# Constants loaded from config
G = config["simulation"]["g"]
L_PENDULUM = config["simulation"]["L"]
TOTAL_TIME = config["simulation"]["total_time"]
NUM_STEPS = config["simulation"]["num_steps"]
TRUE_INITIAL_THETA = config["target_trajectory"]["initial_theta"]
TRUE_INITIAL_OMEGA = config["target_trajectory"]["initial_omega"]
EPOCHS = config["optimization"]["epochs"]
LEARNING_RATE = config["optimization"]["learning_rate"]
SEED = config["optimization"]["seed"]
# Derived constants
DT = TOTAL_TIME / NUM_STEPS
TIME_POINTS = jnp.linspace(0, TOTAL_TIME, NUM_STEPS)