FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Training scripts · geoelements/DPCFunctionEncoder@656d92b · GitHub

Commit 656d92b

Browse files
Training scripts
1 parent 671ae98 commit 656d92b

2 files changed

Lines changed: 297 additions & 0 deletions

File tree

‎src/1.train_function_encoder.py‎

Lines changed: 161 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,161 @@
1+
import argparse
2+
import random
3+
import datetime
4+
import numpy
5+
import torch
6+
import tqdm
7+
import os
8+
9+
from function_encoder.utils.training import train_step
10+
from torch.utils.data import DataLoader
11+
from torch.utils.tensorboard import SummaryWriter
12+
13+
from Datasets.get_dataset import get_function_encoder_dataset
14+
from getters import create_function_encoder
15+
16+
torch.cuda.set_device(1)
17+
torch.set_printoptions(precision=16)
18+
torch.set_default_dtype(torch.float64)
19+
20+
if __name__ == "__main__":
21+
22+
# training arguments
23+
parser = argparse.ArgumentParser(description="Train a model with specified parameters.")
24+
parser.add_argument("--grad_steps", type=int, default=10_000, help="Number of training epochs")
25+
parser.add_argument("--batch_size", type=int, default=32, help="Batch size for training")
26+
parser.add_argument("--seed", type=int, default=0, help="RNG seed")
27+
parser.add_argument("--device", type=str, default="cuda", help="Torch device to use.")
28+
parser.add_argument("--n_basis", type=int, default=11,)
29+
parser.add_argument("--n_layers", type=int, default=4,)
30+
parser.add_argument("--n_hidden", type=int, default=77,)
31+
parser.add_argument("--log_dir", type=str, default="logs/function_encoder", )
32+
parser.add_argument("--dataset", type=str, default="Vanderpol", )
33+
parser.add_argument("--use_residual", type=bool, default=False)
34+
args = parser.parse_args()
35+
36+
# create a logdir
37+
datetime = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
38+
args.log_dir = os.path.join(args.log_dir, args.dataset, f"seed_{args.seed}", datetime)
39+
40+
# create a summary writer
41+
logger = SummaryWriter(args.log_dir)
42+
43+
# Set the random seed for reproducibility
44+
random.seed(args.seed)
45+
numpy.random.seed(args.seed)
46+
torch.manual_seed(args.seed)
47+
48+
# Fetch a dataset
49+
train_dataset, eval_dataset = get_function_encoder_dataset(args)
50+
dataloader = DataLoader(train_dataset, batch_size=args.batch_size)
51+
dataloader_iter = iter(dataloader)
52+
eval_dataloader = DataLoader(eval_dataset, batch_size=args.batch_size)
53+
eval_dataloader_iter = iter(dataloader)
54+
# state_weights = train_dataset.weights.to(args.device)
55+
# initialize the model
56+
model = create_function_encoder(
57+
state_size=train_dataset.state_size,
58+
action_size=train_dataset.action_size,
59+
n_hidden=args.n_hidden,
60+
n_layers=args.n_layers,
61+
n_basis=args.n_basis,
62+
use_residual=args.use_residual,
63+
device=args.device,
64+
)
65+
66+
# MSE loss function
67+
def train_loss_function(model, batch):
68+
69+
_, y0, u0, dt, y1, y0_example, u0_example, dt_example, y1_example = batch
70+
71+
# change device
72+
y0 = y0.to(args.device)
73+
u0 = u0.to(args.device)
74+
dt = dt.to(args.device)
75+
y1 = y1.to(args.device)
76+
y0_example = y0_example.to(args.device)
77+
u0_example = u0_example.to(args.device)
78+
dt_example = dt_example.to(args.device)
79+
y1_example = y1_example.to(args.device)
80+
81+
# compute coefficients
82+
coefficients, _ = model.compute_coefficients((y0_example, u0_example, dt_example), y1_example)
83+
84+
pred = model((y0, u0, dt), coefficients=coefficients)
85+
pred_loss = torch.nn.functional.mse_loss(pred, y1)
86+
87+
# residual loss
88+
if args.use_residual:
89+
residual = model.residual_function((y0, u0, dt))
90+
residual_loss = torch.nn.functional.mse_loss(residual, y1)
91+
else:
92+
residual_loss = torch.tensor(0.0, device=args.device)
93+
94+
return pred_loss + residual_loss
95+
96+
def eval_loss_function(model, batch):
97+
_, y0, u0, dt, y1, y0_example, u0_example, dt_example, y1_example = batch
98+
99+
# change device
100+
y0 = y0.to(args.device)
101+
u0 = u0.to(args.device)
102+
dt = dt.to(args.device)
103+
y1 = y1.to(args.device)
104+
y0_example = y0_example.to(args.device)
105+
u0_example = u0_example.to(args.device)
106+
dt_example = dt_example.to(args.device)
107+
y1_example = y1_example.to(args.device)
108+
109+
# compute coefficients
110+
coefficients, _ = model.compute_coefficients((y0_example, u0_example, dt_example), y1_example)
111+
112+
# # basis function loss
113+
pred = model((y0, u0, dt), coefficients=coefficients)
114+
pred_loss = torch.nn.functional.mse_loss(pred, y1)
115+
return pred_loss
116+
117+
118+
# train the model
119+
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
120+
with tqdm.trange(args.grad_steps) as tqdm_bar:
121+
for epoch in tqdm_bar:
122+
# get data
123+
batch = next(dataloader_iter)
124+
125+
# train
126+
loss = train_step(model, optimizer, batch, train_loss_function)
127+
128+
# eval (MSE only)
129+
with torch.no_grad():
130+
batch = next(eval_dataloader_iter)
131+
eval_loss = eval_loss_function(model, batch)
132+
133+
# log
134+
tqdm_bar.set_postfix_str(f"Loss: {eval_loss:.2e}")
135+
logger.add_scalar("loss/eval", eval_loss, epoch)
136+
logger.add_scalar("loss/train", loss, epoch)
137+
138+
if epoch % 1000 == 0:
139+
# save a checkpoint
140+
checkpoint_path = os.path.join(args.log_dir, f"checkpoint_epoch_{epoch}.pth")
141+
torch.save({
142+
'epoch': epoch,
143+
'model_state_dict': model.state_dict(),
144+
'optimizer_state_dict': optimizer.state_dict(),
145+
'loss': loss,
146+
}, checkpoint_path)
147+
print(f"Checkpoint saved at {checkpoint_path}")
148+
149+
# save the model
150+
arch_params = {"n_basis": args.n_basis,
151+
"n_layers": args.n_layers,
152+
"n_hidden": args.n_hidden,
153+
"state_size": train_dataset.state_size,
154+
"action_size": train_dataset.action_size,
155+
"use_residual": args.use_residual,
156+
}
157+
torch.save(model.state_dict(), os.path.join(args.log_dir, "model.pth"))
158+
torch.save(arch_params, os.path.join(args.log_dir, "arch_params.pth"))
159+
160+
# plot the result.
161+
train_dataset.plot(model, args)

‎src/2.train_policy.py‎

Lines changed: 136 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,136 @@
1+
import argparse
2+
import datetime
3+
import random
4+
import numpy
5+
6+
from Datasets.get_dataset import *
7+
from getters import find_latest, load_function_encoder, get_coefficients
8+
from neuromancer.modules.activations import activations
9+
from neuromancer.system import Node, System
10+
from neuromancer.loss import PenaltyLoss
11+
from neuromancer.problem import Problem
12+
from neuromancer.trainer import Trainer
13+
14+
from Callbacks import TensorboardCallback, ProgressBarCallback, ListCallback, EvalCallback
15+
from Policies.Policy import Policy
16+
from getters import get_policy
17+
18+
if __name__ == "__main__":
19+
20+
# training arguments
21+
parser = argparse.ArgumentParser(description="Train a model with specified parameters.")
22+
parser.add_argument("--batch_size", type=int, default=32, help="Batch size for training")
23+
parser.add_argument("--num_envs", type=int, default=32, help="Number of dynamical systems to train on")
24+
parser.add_argument("--num_epochs", type=int, default=100, help="Number of epochs to train over")
25+
parser.add_argument("--horizon", type=int, default=50, help="Prediction horizon")
26+
parser.add_argument("--seed", type=int, default=0, help="RNG seed")
27+
parser.add_argument("--device", type=str, default="cuda", help="Torch device to use.")
28+
parser.add_argument("--log_dir", type=str, default="logs/policy", )
29+
parser.add_argument("--n_layers", type=int, default=4,)
30+
parser.add_argument("--n_hidden", type=int, default=256, )
31+
parser.add_argument("--fe_load_path", type=str, default="latest", help="Path to the function encoder model. If 'latest', it will find the latest model in the log directory.")
32+
parser.add_argument("--dataset", type=str, default="VanDerPol", )
33+
parser.add_argument("--policy_type", type=str, default="adaptive", )
34+
args = parser.parse_args()
35+
36+
# create a logdir
37+
datetime = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
38+
args.log_dir = os.path.join(args.log_dir, args.dataset, args.policy_type, f"seed_{args.seed}", datetime)
39+
40+
print(f"Training with arguments: {args}")
41+
42+
# Set the random seed for reproducibility
43+
random.seed(args.seed)
44+
numpy.random.seed(args.seed)
45+
torch.manual_seed(args.seed)
46+
torch.set_default_dtype(torch.float64)
47+
48+
# Fetch a dataset
49+
dataset, _ = get_function_encoder_dataset(args)
50+
51+
args.fe_load_path = "logs/function_encoder/Vanderpol/seed_0/2025-10-21_21-08-35"
52+
print(f"Using latest function encoder model from {args.fe_load_path}")
53+
54+
# load a model, to be used as dynamics with no gradients.
55+
model = load_function_encoder(
56+
load_path=args.fe_load_path,
57+
device=args.device,
58+
)
59+
60+
# get a large set of coefficients from the training dataset
61+
coefficients, hp = get_coefficients(dataset, args, model)
62+
63+
# Training dataset generation
64+
train_data, dev_data = get_trajectory_dataset(args, coefficients, hp)
65+
66+
# prepare to train
67+
train_loader = torch.utils.data.DataLoader(train_data, batch_size=args.batch_size, collate_fn=train_data.collate_fn, shuffle=False)
68+
dev_loader = torch.utils.data.DataLoader(dev_data, batch_size=args.batch_size, collate_fn=dev_data.collate_fn, shuffle=False)
69+
70+
71+
# create the learned dynamics model
72+
dt = torch.tensor([[dataset.dt]], device=args.device)
73+
def model_wrapper(x, u, c):
74+
ret = (
75+
# Function encoder expects a leading batch dimension, and a separate dt for every point.
76+
model((x.unsqueeze(1), u.unsqueeze(1), dt.expand(x.shape[0], 1)), coefficients=c)
77+
+ x.unsqueeze(1)
78+
).squeeze(1)
79+
return ret
80+
model_node = Node(model_wrapper, ["x", "u", 'c'], ["x"], name="model")
81+
82+
# create the neural net control policy
83+
policy = get_policy(args, dataset, coefficients, len(model.basis_functions.basis_functions),)
84+
85+
86+
# create the closed loop system
87+
cl_system = System([policy, model_node], nsteps=args.horizon)
88+
objectives, constraints = train_data.get_constraints_objectives(args.device)
89+
components = [cl_system]
90+
loss = PenaltyLoss(objectives, constraints)
91+
problem = Problem(components, loss)
92+
93+
# create callbacks
94+
cb1 = TensorboardCallback(args.log_dir)
95+
cb2 = ProgressBarCallback(len(train_loader) * args.num_epochs)
96+
cb3 = EvalCallback(dataset, dev_data, model, cb1.summary_writer)
97+
callback = ListCallback([cb1, cb2, cb3])
98+
99+
# train the policy
100+
optimizer = torch.optim.AdamW(policy.parameters(), lr=2e-3)
101+
trainer = Trainer(
102+
problem,
103+
train_loader,
104+
dev_loader,
105+
optimizer=optimizer,
106+
epochs=args.num_epochs,
107+
train_metric='train_loss',
108+
warmup=50,
109+
device=args.device,
110+
callback=callback,
111+
# log_dir=args.log_dir,
112+
)
113+
114+
best_model = trainer.train()
115+
116+
# save the policy
117+
os.makedirs(args.log_dir, exist_ok=True)
118+
trainer.model.load_state_dict(best_model)
119+
torch.save(trainer.model.state_dict(), os.path.join(args.log_dir, "policy.pth"))
120+
121+
# simulate a trajectory in the actual system
122+
with torch.no_grad():
123+
dataloader = DataLoader(dataset, batch_size=1)
124+
hidden_parameter, y0, u0, dt, y1, y0_example, u0_example, dt_example, y1_example = next(iter(dataloader))
125+
y0, u0, dt, y1, y0_example, u0_example, dt_example, y1_example = (
126+
y0.to(args.device),
127+
u0.to(args.device),
128+
dt.to(args.device),
129+
y1.to(args.device),
130+
y0_example.to(args.device),
131+
u0_example.to(args.device),
132+
dt_example.to(args.device),
133+
y1_example.to(args.device),
134+
)
135+
coefficients, _ = model.compute_coefficients((y0_example, u0_example, dt_example), y1_example)
136+
dev_data.rollout_real_trajectory(hidden_parameter, coefficients, cl_system.nodes[0], args.log_dir)

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL