(module_list)
| 30 | |
| 31 | # count # of param for a list of module |
| 32 | def count_param(module_list): |
| 33 | return sum(x.numel() for module in module_list for x in module.parameters()) / 10**6 |
| 34 | |
| 35 | # display the peak memory of cuda |
| 36 | def print_peak_memory(prefix, device): |