(args, opts)
| 206 | optimizer.step() |
| 207 | |
| 208 | def train_with_config(args, opts): |
| 209 | print(args) |
| 210 | try: |
| 211 | os.makedirs(opts.checkpoint) |
| 212 | except OSError as e: |
| 213 | if e.errno != errno.EEXIST: |
| 214 | raise RuntimeError('Unable to create checkpoint directory:', opts.checkpoint) |
| 215 | train_writer = tensorboardX.SummaryWriter(os.path.join(opts.checkpoint, "logs")) |
| 216 | |
| 217 | |
| 218 | print('Loading dataset...') |
| 219 | trainloader_params = { |
| 220 | 'batch_size': args.batch_size, |
| 221 | 'shuffle': True, |
| 222 | 'num_workers': 12, |
| 223 | 'pin_memory': True, |
| 224 | 'prefetch_factor': 4, |
| 225 | 'persistent_workers': True |
| 226 | } |
| 227 | |
| 228 | testloader_params = { |
| 229 | 'batch_size': args.batch_size, |
| 230 | 'shuffle': False, |
| 231 | 'num_workers': 12, |
| 232 | 'pin_memory': True, |
| 233 | 'prefetch_factor': 4, |
| 234 | 'persistent_workers': True |
| 235 | } |
| 236 | |
| 237 | train_dataset = MotionDataset3D(args, args.subset_list, 'train') |
| 238 | test_dataset = MotionDataset3D(args, args.subset_list, 'test') |
| 239 | train_loader_3d = DataLoader(train_dataset, **trainloader_params) |
| 240 | test_loader = DataLoader(test_dataset, **testloader_params) |
| 241 | |
| 242 | if args.train_2d: |
| 243 | posetrack = PoseTrackDataset2D() |
| 244 | posetrack_loader_2d = DataLoader(posetrack, **trainloader_params) |
| 245 | instav = InstaVDataset2D() |
| 246 | instav_loader_2d = DataLoader(instav, **trainloader_params) |
| 247 | |
| 248 | datareader = DataReaderH36M(n_frames=args.clip_len, sample_stride=args.sample_stride, data_stride_train=args.data_stride, data_stride_test=args.clip_len, dt_root = 'data/motion3d', dt_file=args.dt_file) |
| 249 | min_loss = 100000 |
| 250 | model_backbone = load_backbone(args) |
| 251 | model_params = 0 |
| 252 | for parameter in model_backbone.parameters(): |
| 253 | model_params = model_params + parameter.numel() |
| 254 | print('INFO: Trainable parameter count:', model_params) |
| 255 | |
| 256 | if torch.cuda.is_available(): |
| 257 | model_backbone = nn.DataParallel(model_backbone) |
| 258 | model_backbone = model_backbone.cuda() |
| 259 | |
| 260 | if args.finetune: |
| 261 | if opts.resume or opts.evaluate: |
| 262 | chk_filename = opts.evaluate if opts.evaluate else opts.resume |
| 263 | print('Loading checkpoint', chk_filename) |
| 264 | checkpoint = torch.load(chk_filename, map_location=lambda storage, loc: storage) |
| 265 | model_backbone.load_state_dict(checkpoint['model_pos'], strict=True) |
no test coverage detected