| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 671ae98 commit 656d92b
2 files changed
| Original file line number | Diff line number | Diff 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) | ||
| Original file line number | Diff line number | Diff 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) | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments