| 298 | return mean_mask |
| 299 | |
| 300 | def cal_params_flops(model, size, logger): |
| 301 | input = torch.randn(1, 3, size, size).cuda() |
| 302 | flops, params = profile(model, inputs=(input,)) |
| 303 | print('flops',flops/1e9) ## 打印计算量 |
| 304 | print('params',params/1e6) ## 打印参数量 |
| 305 | |
| 306 | total = sum(p.numel() for p in model.parameters()) |
| 307 | print("Total params: %.2fM" % (total/1e6)) |
| 308 | logger.info(f'flops: {flops/1e9}, params: {params/1e6}, Total params: : {total/1e6:.4f}') |
| 309 | |
| 310 | # Example function to calculate and print GMACs and parameter count for a given model |
| 311 | def print_model_stats(model, input_size=(3, 224, 224)): |