MCPcopy Create free account
hub / github.com/Walter0807/MotionBERT / train_with_config

Function train_with_config

train.py:208–377  ·  view source on GitHub ↗
(args, opts)

Source from the content-addressed store, hash-verified

206 optimizer.step()
207
208def 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)

Callers 1

train.pyFile · 0.70

Calls 11

MotionDataset3DClass · 0.90
PoseTrackDataset2DClass · 0.90
InstaVDataset2DClass · 0.90
DataReaderH36MClass · 0.90
Augmenter2DClass · 0.90
load_backboneFunction · 0.85
partial_train_layersFunction · 0.85
AverageMeterClass · 0.85
evaluateFunction · 0.85
save_checkpointFunction · 0.85
train_epochFunction · 0.70

Tested by

no test coverage detected