visualize feature map
(args, dl, model, hw)
| 22 | |
| 23 | @torch.no_grad() |
| 24 | def vis_featuremap(args, dl, model, hw): |
| 25 | """ visualize feature map """ |
| 26 | model.eval() |
| 27 | for i, (data, pose, _, _) in enumerate(tqdm(dl, desc="Visualizing Feature Map")): |
| 28 | data = data.to(args.device) |
| 29 | inputs = data.to(args.device) |
| 30 | with autocast(args.device, enabled=args.amp, dtype=args.amp_dtype): |
| 31 | feature_list, predicted_pose = model(inputs, return_feature=True, is_single_stream=True) |
| 32 | feature_target = feature_list[0] |
| 33 | save_feature_gs(args, i, feature_target, hw) |
| 34 | |
| 35 | |
| 36 | def freeze_bn_layer(model): |
no test coverage detected