| 824 | |
| 825 | |
| 826 | def see_memory_usage(message, force=False): |
| 827 | if not force: |
| 828 | return |
| 829 | if dist.is_initialized() and not dist.get_rank() == 0: |
| 830 | return |
| 831 | |
| 832 | # python doesn't do real-time garbage collection so do it explicitly to get the correct RAM reports |
| 833 | gc.collect() |
| 834 | |
| 835 | # Print message except when distributed but not rank 0 |
| 836 | print(message) |
| 837 | print(f"MA {round(get_accelerator().memory_allocated() / (1024 * 1024 * 1024),2 )} GB \ |
| 838 | Max_MA {round(get_accelerator().max_memory_allocated() / (1024 * 1024 * 1024),2)} GB \ |
| 839 | CA {round(torch_memory_reserved() / (1024 * 1024 * 1024),2)} GB \ |
| 840 | Max_CA {round(torch_max_memory_reserved() / (1024 * 1024 * 1024))} GB ") |
| 841 | |
| 842 | vm_stats = psutil.virtual_memory() |
| 843 | used_GB = round(((vm_stats.total - vm_stats.available) / (1024**3)), 2) |
| 844 | print(f'CPU Virtual Memory: used = {used_GB} GB, percent = {vm_stats.percent}%') |
| 845 | |
| 846 | # get the peak memory to report correct data, so reset the counter for the next call |
| 847 | get_accelerator().reset_peak_memory_stats() |
| 848 | |
| 849 | |
| 850 | def call_to_str(base, *args, **kwargs): |