|
| 1 | +import os |
| 2 | +import argparse |
| 3 | +import torch |
| 4 | +from models.st_gcn.st_gcn import STGCNEmbedding |
| 5 | +import models.ResGCNv1 |
| 6 | + |
| 7 | + |
| 8 | +def parse_option(): |
| 9 | + parser = argparse.ArgumentParser(description="Training model on gait sequence") |
| 10 | + parser.add_argument("dataset", choices=["casia-b", "outdoor-gait", "tum-gaid"]) |
| 11 | + parser.add_argument("train_data_path", help="Path to train data CSV") |
| 12 | + parser.add_argument("--valid_data_path", help="Path to validation data CSV") |
| 13 | + parser.add_argument("--valid_split", type=float, default=0.2) |
| 14 | + |
| 15 | + parser.add_argument("--checkpoint_path", help="Path to checkpoint to resume") |
| 16 | + parser.add_argument("--weight_path", help="Path to weights for model") |
| 17 | + |
| 18 | + # Optionals |
| 19 | + parser.add_argument("--num_workers", type=int, default=8) |
| 20 | + parser.add_argument( |
| 21 | + "--gpus", default="0", help="-1 for CPU, use comma for multiple gpus" |
| 22 | + ) |
| 23 | + parser.add_argument("--batch_size", type=int, default=64) |
| 24 | + parser.add_argument("--batch_size_validation", type=int, default=64) |
| 25 | + parser.add_argument("--epochs", type=int, default=500) |
| 26 | + parser.add_argument("--start_epoch", type=int, default=1) |
| 27 | + parser.add_argument("--log_interval", type=int, default=10) |
| 28 | + parser.add_argument("--save_interval", type=int, default=50, help="save frequency") |
| 29 | + parser.add_argument( |
| 30 | + "--save_best_start", type=float, default=0.3, help="save frequency" |
| 31 | + ) |
| 32 | + parser.add_argument("--use_amp", action="store_true") |
| 33 | + parser.add_argument("--tune", action="store_true") |
| 34 | + parser.add_argument("--shuffle", action="store_true") |
| 35 | + parser.add_argument("--exp_name", help="Name of the experiment") |
| 36 | + |
| 37 | + parser.add_argument("--network_name", default="resgcn-n39-r4") |
| 38 | + parser.add_argument("--sequence_length", type=int, default=60) |
| 39 | + parser.add_argument("--embedding_layer_size", type=int, default=256) |
| 40 | + parser.add_argument("--temporal_kernel_size", type=int, default=9) |
| 41 | + parser.add_argument("--dropout", type=float, default=0.4) |
| 42 | + parser.add_argument("--learning_rate", type=float, default=1e-3) |
| 43 | + parser.add_argument( |
| 44 | + "--lr_decay_rate", type=float, default=0.1, help="decay rate for learning rate" |
| 45 | + ) |
| 46 | + parser.add_argument("--point_noise_std", type=float, default=0.05) |
| 47 | + parser.add_argument("--joint_noise_std", type=float, default=0.1) |
| 48 | + parser.add_argument("--flip_probability", type=float, default=0.5) |
| 49 | + parser.add_argument("--mirror_probability", type=float, default=0.5) |
| 50 | + parser.add_argument("--weight_decay", type=float, default=1e-5) |
| 51 | + parser.add_argument("--use_multi_branch", action="store_true") |
| 52 | + parser.add_argument( |
| 53 | + "--temp", type=float, default=0.07, help="temperature for loss function" |
| 54 | + ) |
| 55 | + opt = parser.parse_args() |
| 56 | + |
| 57 | + # Sanitize opts |
| 58 | + opt.gpus_str = opt.gpus |
| 59 | + opt.gpus = [int(gpu) for gpu in opt.gpus.split(",")] |
| 60 | + |
| 61 | + return opt |
| 62 | + |
| 63 | + |
| 64 | +def log_hyperparameter(writer, opt, accuracy, loss): |
| 65 | + writer.add_hparams( |
| 66 | + { |
| 67 | + "batch_size": opt.batch_size, |
| 68 | + "sequence_length": opt.sequence_length, |
| 69 | + "embedding_layer_size": opt.embedding_layer_size, |
| 70 | + "dropout": opt.dropout, |
| 71 | + "learning_rate": opt.learning_rate, |
| 72 | + "lr_decay_rate": opt.lr_decay_rate, |
| 73 | + "point_noise_std": opt.point_noise_std, |
| 74 | + "weight_decay": opt.weight_decay, |
| 75 | + "temp": opt.temp, |
| 76 | + }, |
| 77 | + { |
| 78 | + "hparam/accuracy": accuracy, |
| 79 | + "hparam/loss": loss, |
| 80 | + }, |
| 81 | + ) |
| 82 | + |
| 83 | + |
| 84 | +def setup_environment(opt): |
| 85 | + # HACK: Fix tensorboard |
| 86 | + import tensorflow as tf |
| 87 | + import tensorboard as tb |
| 88 | + |
| 89 | + tf.io.gfile = tb.compat.tensorflow_stub.io.gfile |
| 90 | + |
| 91 | + os.environ["CUDA_VISIBLE_DEVICES"] = opt.gpus_str |
| 92 | + opt.cuda = opt.gpus[0] >= 0 |
| 93 | + torch.device("cuda" if opt.cuda else "cpu") |
| 94 | + |
| 95 | + return opt |
| 96 | + |
| 97 | + |
| 98 | +def get_model_stgcn(opt): |
| 99 | + # Model |
| 100 | + input_channels = 3 |
| 101 | + edge_importance_weighting = True |
| 102 | + graph_args = {"strategy": "spatial"} |
| 103 | + |
| 104 | + embedding_net = STGCNEmbedding( |
| 105 | + input_channels, |
| 106 | + graph_args, |
| 107 | + edge_importance_weighting=edge_importance_weighting, |
| 108 | + embedding_layer_size=opt.embedding_layer_size, |
| 109 | + temporal_kernel_size=opt.temporal_kernel_size, |
| 110 | + dropout=opt.dropout, |
| 111 | + ) |
| 112 | + |
| 113 | + return embedding_net |
| 114 | + |
| 115 | + |
| 116 | +def get_model_resgcn(graph, opt): |
| 117 | + model_args = { |
| 118 | + "A": torch.tensor(graph.A, dtype=torch.float32, requires_grad=False), |
| 119 | + "num_class": opt.embedding_layer_size, |
| 120 | + "num_input": 1 if not opt.use_multi_branch else 3, |
| 121 | + "num_channel": 3 if not opt.use_multi_branch else 6, |
| 122 | + "parts": graph.parts, |
| 123 | + } |
| 124 | + return models.ResGCNv1.create(opt.network_name, **model_args) |
| 125 | + |
| 126 | + |
| 127 | +def get_trainer(model, opt, steps_per_epoch): |
| 128 | + optimizer = torch.optim.Adam( |
| 129 | + model.parameters(), lr=opt.learning_rate, weight_decay=opt.weight_decay |
| 130 | + ) |
| 131 | + scheduler = torch.optim.lr_scheduler.OneCycleLR( |
| 132 | + optimizer, opt.learning_rate, epochs=opt.epochs, steps_per_epoch=steps_per_epoch |
| 133 | + ) |
| 134 | + scaler = torch.cuda.amp.GradScaler(enabled=opt.use_amp) |
| 135 | + |
| 136 | + return optimizer, scheduler, scaler |
| 137 | + |
| 138 | + |
| 139 | +def load_checkpoint(model, optimizer, scheduler, scaler, opt): |
| 140 | + if opt.checkpoint_path is not None: |
| 141 | + checkpoint = torch.load(opt.checkpoint_path) |
| 142 | + model.load_state_dict(checkpoint["model"]) |
| 143 | + optimizer.load_state_dict(checkpoint["optimizer"]) |
| 144 | + scheduler.load_state_dict(checkpoint["scheduler"]) |
| 145 | + scaler.load_state_dict(checkpoint["scaler"]) |
| 146 | + opt.start_epoch = checkpoint["epoch"] |
| 147 | + |
| 148 | + if opt.weight_path is not None: |
| 149 | + checkpoint = torch.load(opt.weight_path) |
| 150 | + model.load_state_dict(checkpoint["model"], strict=False) |
| 151 | + |
| 152 | + |
| 153 | +def save_model(model, optimizer, scheduler, scaler, opt, epoch, save_file): |
| 154 | + print("==> Saving...") |
| 155 | + state = { |
| 156 | + "opt": opt, |
| 157 | + "model": model.state_dict(), |
| 158 | + "optimizer": optimizer.state_dict(), |
| 159 | + "scheduler": scheduler.state_dict(), |
| 160 | + "scaler": scaler.state_dict(), |
| 161 | + "epoch": epoch, |
| 162 | + } |
| 163 | + torch.save(state, save_file) |
| 164 | + del state |
| 165 | + |
| 166 | + |
| 167 | +def count_parameters(model): |
| 168 | + """ |
| 169 | + Useful function to compute number of parameters in a model. |
| 170 | + """ |
| 171 | + return sum(p.numel() for p in model.parameters() if p.requires_grad) |
0 commit comments