[lint] Update lint (#16)
* update lint * update readme * update ruff lint
This commit is contained in:
+56
-37
@@ -4,23 +4,28 @@ from torch.cuda import nvtx
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List
|
||||
|
||||
|
||||
def _is_distributed():
|
||||
return torch.distributed.is_available() and torch.distributed.is_initialized()
|
||||
|
||||
|
||||
def _get_world_size():
|
||||
if _is_distributed():
|
||||
return torch.distributed.get_world_size()
|
||||
return 1
|
||||
|
||||
|
||||
def _get_rank():
|
||||
if _is_distributed():
|
||||
return torch.distributed.get_rank()
|
||||
return 0
|
||||
|
||||
|
||||
def _barrier(group=None):
|
||||
if _is_distributed():
|
||||
torch.distributed.barrier(group=group)
|
||||
|
||||
|
||||
if torch.distributed.is_available():
|
||||
try:
|
||||
dist_all_gather_func = torch.distributed.all_gather_into_tensor
|
||||
@@ -29,6 +34,7 @@ if torch.distributed.is_available():
|
||||
else:
|
||||
dist_all_gather_func = None
|
||||
|
||||
|
||||
class TimerBase(ABC):
|
||||
"""Timer base class."""
|
||||
|
||||
@@ -76,7 +82,7 @@ class DummyTimer(TimerBase):
|
||||
"""Dummy Timer."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__('dummy timer')
|
||||
super().__init__("dummy timer")
|
||||
|
||||
def start(self, barrier=False, nvtx_push=False):
|
||||
return
|
||||
@@ -89,8 +95,8 @@ class DummyTimer(TimerBase):
|
||||
|
||||
def elapsed(self, reset=True, barrier=False):
|
||||
raise Exception(
|
||||
'dummy timer should not be used to calculate elapsed time, '
|
||||
'check if timer\'s log_level <= self._log_level.'
|
||||
"dummy timer should not be used to calculate elapsed time, "
|
||||
"check if timer's log_level <= self._log_level."
|
||||
)
|
||||
|
||||
def active_time(self):
|
||||
@@ -98,8 +104,8 @@ class DummyTimer(TimerBase):
|
||||
Note: Not supported for DummyTimer.
|
||||
"""
|
||||
raise Exception(
|
||||
'active timer should not be used to calculate elapsed time, '
|
||||
'check if timer\'s log_level <= self._log_level.'
|
||||
"active timer should not be used to calculate elapsed time, "
|
||||
"check if timer's log_level <= self._log_level."
|
||||
)
|
||||
|
||||
|
||||
@@ -144,7 +150,7 @@ class Timer(TimerBase):
|
||||
Args:
|
||||
barrier (bool, optional): Synchronizes ranks before starting. Defaults to False.
|
||||
"""
|
||||
assert not self._started, 'timer has already been started'
|
||||
assert not self._started, "timer has already been started"
|
||||
if barrier:
|
||||
_barrier(group=self._barrier_group)
|
||||
if torch.cuda.is_available():
|
||||
@@ -154,7 +160,6 @@ class Timer(TimerBase):
|
||||
if nvtx_push:
|
||||
nvtx.range_push("{}".format(self.name))
|
||||
self.nvtx = True
|
||||
|
||||
|
||||
def stop(self, barrier=False, sync=False):
|
||||
"""Stop the timer.
|
||||
@@ -164,7 +169,7 @@ class Timer(TimerBase):
|
||||
"""
|
||||
if self.nvtx:
|
||||
nvtx.range_pop()
|
||||
assert self._started, 'timer is not started'
|
||||
assert self._started, "timer is not started"
|
||||
if barrier:
|
||||
_barrier(group=self._barrier_group)
|
||||
if torch.cuda.is_available() and sync:
|
||||
@@ -221,10 +226,10 @@ class Timers:
|
||||
Allowed: ['max', 'minmax', 'all'].
|
||||
"""
|
||||
self._log_level = log_level
|
||||
allowed_log_options = set(['max', 'minmax', 'all'])
|
||||
allowed_log_options = set(["max", "minmax", "all"])
|
||||
assert (
|
||||
log_option in allowed_log_options
|
||||
), 'input log option {} is invalid. It must be one of {}'.format(
|
||||
), "input log option {} is invalid. It must be one of {}".format(
|
||||
log_option, allowed_log_options
|
||||
)
|
||||
self._log_option = log_option
|
||||
@@ -240,8 +245,10 @@ class Timers:
|
||||
if name in self._timers:
|
||||
if log_level is not None:
|
||||
assert log_level == self._log_levels[name], (
|
||||
'input log level {} does not match already existing '
|
||||
'log level {} for {} timer'.format(log_level, self._log_levels[name], name)
|
||||
"input log level {} does not match already existing "
|
||||
"log level {} for {} timer".format(
|
||||
log_level, self._log_levels[name], name
|
||||
)
|
||||
)
|
||||
return self._timers[name]
|
||||
# If timer does not exist and no log level is provided,
|
||||
@@ -250,7 +257,7 @@ class Timers:
|
||||
log_level = self._max_log_level
|
||||
assert (
|
||||
log_level <= self._max_log_level
|
||||
), 'log level {} is larger than max supported log level {}'.format(
|
||||
), "log level {} is larger than max supported log level {}".format(
|
||||
log_level, self._max_log_level
|
||||
)
|
||||
# Now if the input log level is larger than the one set for
|
||||
@@ -284,7 +291,7 @@ class Timers:
|
||||
if torch.cuda.is_available():
|
||||
device = torch.cuda.current_device()
|
||||
else:
|
||||
device = torch.device('cpu')
|
||||
device = torch.device("cpu")
|
||||
|
||||
rank_name_to_time = torch.zeros(
|
||||
(world_size, len(names)), dtype=torch.float, device=device
|
||||
@@ -296,7 +303,9 @@ class Timers:
|
||||
|
||||
if world_size > 1 and _is_distributed() and dist_all_gather_func is not None:
|
||||
try:
|
||||
dist_all_gather_func(rank_name_to_time.view(-1), rank_name_to_time[rank, :].view(-1))
|
||||
dist_all_gather_func(
|
||||
rank_name_to_time.view(-1), rank_name_to_time[rank, :].view(-1)
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Warning: all_gather failed: {e}. Using single rank timing.")
|
||||
|
||||
@@ -319,30 +328,38 @@ class Timers:
|
||||
)
|
||||
return name_to_min_max_time
|
||||
|
||||
def _get_global_min_max_time_string(self, names, reset, barrier, normalizer, max_only):
|
||||
def _get_global_min_max_time_string(
|
||||
self, names, reset, barrier, normalizer, max_only
|
||||
):
|
||||
"""Report strings for max/minmax times across all ranks."""
|
||||
name_to_min_max_time = self._get_global_min_max_time(names, reset, barrier, normalizer)
|
||||
name_to_min_max_time = self._get_global_min_max_time(
|
||||
names, reset, barrier, normalizer
|
||||
)
|
||||
if not name_to_min_max_time:
|
||||
return None
|
||||
|
||||
|
||||
world_size = _get_world_size()
|
||||
if world_size == 1:
|
||||
output_string = 'time (ms):'
|
||||
output_string = "time (ms):"
|
||||
for name in name_to_min_max_time:
|
||||
_, max_time = name_to_min_max_time[name]
|
||||
output_string += '\n {}: {:.2f}'.format((name + ' ').ljust(48, '.'), max_time)
|
||||
output_string += "\n {}: {:.2f}".format(
|
||||
(name + " ").ljust(48, "."), max_time
|
||||
)
|
||||
else:
|
||||
if max_only:
|
||||
output_string = 'max time across ranks (ms):'
|
||||
output_string = "max time across ranks (ms):"
|
||||
else:
|
||||
output_string = '(min, max) time across ranks (ms):'
|
||||
output_string = "(min, max) time across ranks (ms):"
|
||||
for name in name_to_min_max_time:
|
||||
min_time, max_time = name_to_min_max_time[name]
|
||||
if max_only:
|
||||
output_string += '\n {}: {:.2f}'.format((name + ' ').ljust(48, '.'), max_time)
|
||||
output_string += "\n {}: {:.2f}".format(
|
||||
(name + " ").ljust(48, "."), max_time
|
||||
)
|
||||
else:
|
||||
output_string += '\n {}: ({:.2f}, {:.2f})'.format(
|
||||
(name + ' ').ljust(48, '.'), min_time, max_time
|
||||
output_string += "\n {}: ({:.2f}, {:.2f})".format(
|
||||
(name + " ").ljust(48, "."), min_time, max_time
|
||||
)
|
||||
return output_string
|
||||
|
||||
@@ -351,7 +368,7 @@ class Timers:
|
||||
rank_name_to_time = self._get_elapsed_time_all_ranks(names, reset, barrier)
|
||||
world_size = _get_world_size()
|
||||
|
||||
output_string = 'times across ranks (ms):'
|
||||
output_string = "times across ranks (ms):"
|
||||
no_reported_timing = True
|
||||
for i, name in enumerate(names):
|
||||
not_yet_found = True
|
||||
@@ -360,13 +377,13 @@ class Timers:
|
||||
no_reported_timing = False
|
||||
if not_yet_found:
|
||||
not_yet_found = False
|
||||
output_string += '\n {}:'.format(name)
|
||||
output_string += "\n {}:".format(name)
|
||||
if world_size == 1:
|
||||
output_string += '\n {:.2f}'.format(
|
||||
output_string += "\n {:.2f}".format(
|
||||
rank_name_to_time[rank, i] / normalizer
|
||||
)
|
||||
else:
|
||||
output_string += '\n rank {:2d}: {:.2f}'.format(
|
||||
output_string += "\n rank {:2d}: {:.2f}".format(
|
||||
rank, rank_name_to_time[rank, i] / normalizer
|
||||
)
|
||||
if no_reported_timing:
|
||||
@@ -398,23 +415,23 @@ class Timers:
|
||||
str: Formatted string with the timer values.
|
||||
"""
|
||||
|
||||
if names == None: # get all registered timers
|
||||
if names is None: # get all registered timers
|
||||
names = list(self._timers.keys())
|
||||
|
||||
assert normalizer > 0.0
|
||||
if self._log_option in ['max', 'minmax']:
|
||||
if self._log_option in ["max", "minmax"]:
|
||||
max_only = False
|
||||
if self._log_option == 'max':
|
||||
if self._log_option == "max":
|
||||
max_only = True
|
||||
output_string = self._get_global_min_max_time_string(
|
||||
names, reset, barrier, normalizer / 1000.0, max_only
|
||||
)
|
||||
elif self._log_option == 'all':
|
||||
elif self._log_option == "all":
|
||||
output_string = self._get_all_ranks_time_string(
|
||||
names, reset, barrier, normalizer / 1000.0
|
||||
)
|
||||
else:
|
||||
raise Exception('unknown timing log option {}'.format(self._log_option))
|
||||
raise Exception("unknown timing log option {}".format(self._log_option))
|
||||
return output_string
|
||||
|
||||
def log(
|
||||
@@ -444,7 +461,7 @@ class Timers:
|
||||
# If no input rank is provided, log on last rank.
|
||||
world_size = _get_world_size()
|
||||
current_rank = _get_rank()
|
||||
|
||||
|
||||
if rank is None:
|
||||
rank = world_size - 1
|
||||
if rank == current_rank and output_string is not None:
|
||||
@@ -476,8 +493,10 @@ class Timers:
|
||||
# torch.utils.add_scalars makes each timer its own run, which
|
||||
# polutes the runs list, so we just add each as a scalar
|
||||
assert normalizer > 0.0
|
||||
name_to_min_max_time = self._get_global_min_max_time(names, reset, barrier, normalizer)
|
||||
name_to_min_max_time = self._get_global_min_max_time(
|
||||
names, reset, barrier, normalizer
|
||||
)
|
||||
if writer is not None:
|
||||
for name in name_to_min_max_time:
|
||||
_, max_time = name_to_min_max_time[name]
|
||||
writer.add_scalar(name + '-time', max_time, iteration)
|
||||
writer.add_scalar(name + "-time", max_time, iteration)
|
||||
|
||||
Reference in New Issue
Block a user