Files
VLA/wall_x/utils/timers.py
T

441 lines
14 KiB
Python
Raw Normal View History

2026-06-15 11:40:00 +08:00
import logging
import os
2025-09-07 14:59:17 +08:00
import time
from abc import ABC, abstractmethod
2026-06-15 11:40:00 +08:00
from contextlib import nullcontext
from functools import wraps
2025-09-07 14:59:17 +08:00
from typing import List
2026-02-03 11:35:25 +08:00
import torch
2026-06-15 11:40:00 +08:00
from torch.cuda import nvtx
logger = logging.getLogger(__name__)
2025-09-11 13:18:33 +08:00
2026-02-03 11:35:25 +08:00
ENABLE_PERFORMANCE_TIMING = (
os.environ.get("ENABLE_PERFORMANCE_TIMING", "True").lower() == "true"
)
ENABLE_CUDA_SYNC_IN_TIMER = (
os.environ.get("ENABLE_CUDA_SYNC_IN_TIMER", "False").lower() == "true"
)
class ScopeTimerContext:
def __init__(self, msg):
self.msg = msg
def __enter__(self):
if ENABLE_CUDA_SYNC_IN_TIMER and torch.cuda.is_available():
torch.cuda.synchronize()
self.start_time = time.perf_counter()
return self
def __exit__(self, exc_type, exc_value, traceback):
if ENABLE_CUDA_SYNC_IN_TIMER and torch.cuda.is_available():
torch.cuda.synchronize()
end_time = time.perf_counter()
cost_ms = (end_time - self.start_time) * 1e3
2026-06-15 11:40:00 +08:00
logger.info("%s took %.3f ms to execute", self.msg, cost_ms)
2026-02-03 11:35:25 +08:00
ScopeTimer = ScopeTimerContext if ENABLE_PERFORMANCE_TIMING else nullcontext
def timer(func, msg=None):
2026-06-15 11:40:00 +08:00
"""Decorator to measure function execution time."""
2026-02-03 11:35:25 +08:00
if msg is None:
msg = func.__name__
else:
msg = f"{func.__name__:} {msg}"
@wraps(func)
def wrapper(*args, **kwargs):
with ScopeTimer(msg):
result = func(*args, **kwargs)
return result
return wrapper
2025-09-07 14:59:17 +08:00
def _is_distributed():
return torch.distributed.is_available() and torch.distributed.is_initialized()
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def _get_world_size():
if _is_distributed():
return torch.distributed.get_world_size()
return 1
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def _get_rank():
if _is_distributed():
return torch.distributed.get_rank()
return 0
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def _barrier(group=None):
if _is_distributed():
torch.distributed.barrier(group=group)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
if torch.distributed.is_available():
try:
dist_all_gather_func = torch.distributed.all_gather_into_tensor
except AttributeError:
dist_all_gather_func = torch.distributed.all_gather
else:
dist_all_gather_func = None
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
class TimerBase(ABC):
"""Timer base class."""
def __init__(self, name):
self.name = name
@abstractmethod
def start(self, barrier=False):
2026-06-15 11:40:00 +08:00
"""Start the timer, optionally syncing all ranks with a barrier first."""
2025-09-07 14:59:17 +08:00
pass
@abstractmethod
def stop(self, barrier=False):
2026-06-15 11:40:00 +08:00
"""Stop the timer, optionally syncing all ranks with a barrier first."""
2025-09-07 14:59:17 +08:00
pass
@abstractmethod
def reset(self):
2026-06-15 11:40:00 +08:00
"""Reset accumulated elapsed time to zero."""
2025-09-07 14:59:17 +08:00
pass
@abstractmethod
def elapsed(self, reset=True, barrier=False):
2026-06-15 11:40:00 +08:00
"""Return accumulated elapsed time in seconds; reset if reset=True."""
2025-09-07 14:59:17 +08:00
pass
class DummyTimer(TimerBase):
2026-06-15 11:40:00 +08:00
"""Dummy Timer - no-op placeholder used when log level exceeds threshold."""
2025-09-07 14:59:17 +08:00
def __init__(self):
2025-09-11 13:18:33 +08:00
super().__init__("dummy timer")
2025-09-07 14:59:17 +08:00
2026-06-15 11:40:00 +08:00
def start(self, barrier=False, nvtx_push=False, sync=False, **kwargs):
2025-09-07 14:59:17 +08:00
return
2026-06-15 11:40:00 +08:00
def stop(self, barrier=False, sync=False, **kwargs):
2025-09-07 14:59:17 +08:00
return
def reset(self):
return
def elapsed(self, reset=True, barrier=False):
raise Exception(
2025-09-11 13:18:33 +08:00
"dummy timer should not be used to calculate elapsed time, "
"check if timer's log_level <= self._log_level."
2025-09-07 14:59:17 +08:00
)
def active_time(self):
raise Exception(
2025-09-11 13:18:33 +08:00
"active timer should not be used to calculate elapsed time, "
"check if timer's log_level <= self._log_level."
2025-09-07 14:59:17 +08:00
)
class Timer(TimerBase):
"""
Timer class with ability to start/stop.
Comment on using `barrier`: If this flag is passed, then all
the caller processes will wait till all reach the timing routine.
It is up to the user to make sure all the ranks in `barrier_group`
call it otherwise, it will result in a hang.
Comment on `barrier_group`: By default it is set to None which
in torch distributed land, it will result in the global communicator.
"""
def __init__(self, name):
super().__init__(name)
self._elapsed = 0.0
self._active_time = 0.0
self._started = False
self._barrier_group = None
self._start_time = time.time()
self.nvtx = False
def set_barrier_group(self, barrier_group):
self._barrier_group = barrier_group
2026-02-03 11:35:25 +08:00
def start(self, barrier=False, nvtx_push=False, sync=False):
2025-09-11 13:18:33 +08:00
assert not self._started, "timer has already been started"
2025-09-07 14:59:17 +08:00
if barrier:
_barrier(group=self._barrier_group)
2026-02-03 11:35:25 +08:00
if torch.cuda.is_available() and sync:
2025-09-07 14:59:17 +08:00
torch.cuda.synchronize()
self._start_time = time.time()
self._started = True
if nvtx_push:
nvtx.range_push("{}".format(self.name))
self.nvtx = True
def stop(self, barrier=False, sync=False):
if self.nvtx:
nvtx.range_pop()
2025-09-11 13:18:33 +08:00
assert self._started, "timer is not started"
2025-09-07 14:59:17 +08:00
if barrier:
_barrier(group=self._barrier_group)
if torch.cuda.is_available() and sync:
torch.cuda.synchronize()
elapsed = time.time() - self._start_time
self._elapsed += elapsed
self._active_time += elapsed
self._started = False
def reset(self):
self._elapsed = 0.0
self._started = False
def elapsed(self, reset=True, barrier=False):
_started = self._started
if self._started:
self.stop(barrier=barrier)
_elapsed = self._elapsed
if reset:
self.reset()
if _started:
self.start(barrier=barrier)
return _elapsed
def active_time(self):
return self._active_time
class Timers:
"""Class for a group of Timers."""
def __init__(self, log_level, log_option):
self._log_level = log_level
2025-09-11 13:18:33 +08:00
allowed_log_options = set(["max", "minmax", "all"])
2025-09-07 14:59:17 +08:00
assert (
log_option in allowed_log_options
2025-09-11 13:18:33 +08:00
), "input log option {} is invalid. It must be one of {}".format(
2025-09-07 14:59:17 +08:00
log_option, allowed_log_options
)
self._log_option = log_option
self._timers = {}
self._log_levels = {}
self._dummy_timer = DummyTimer()
self._max_log_level = 2
def __call__(self, name, log_level=None):
if name in self._timers:
if log_level is not None:
assert log_level == self._log_levels[name], (
2025-09-11 13:18:33 +08:00
"input log level {} does not match already existing "
"log level {} for {} timer".format(
log_level, self._log_levels[name], name
)
2025-09-07 14:59:17 +08:00
)
return self._timers[name]
if log_level is None:
log_level = self._max_log_level
assert (
log_level <= self._max_log_level
2025-09-11 13:18:33 +08:00
), "log level {} is larger than max supported log level {}".format(
2025-09-07 14:59:17 +08:00
log_level, self._max_log_level
)
if log_level > self._log_level:
return self._dummy_timer
self._timers[name] = Timer(name)
self._log_levels[name] = log_level
return self._timers[name]
def _get_elapsed_time_all_ranks(self, names, reset, barrier):
if barrier:
_barrier()
world_size = _get_world_size()
rank = _get_rank()
if torch.cuda.is_available():
device = torch.cuda.current_device()
else:
2025-09-11 13:18:33 +08:00
device = torch.device("cpu")
2025-09-07 14:59:17 +08:00
rank_name_to_time = torch.zeros(
(world_size, len(names)), dtype=torch.float, device=device
)
for i, name in enumerate(names):
if name in self._timers:
rank_name_to_time[rank, i] = self._timers[name].elapsed(reset=reset)
if world_size > 1 and _is_distributed() and dist_all_gather_func is not None:
try:
2025-09-11 13:18:33 +08:00
dist_all_gather_func(
rank_name_to_time.view(-1), rank_name_to_time[rank, :].view(-1)
)
2025-09-07 14:59:17 +08:00
except Exception as e:
2026-06-15 11:40:00 +08:00
logger.warning("all_gather failed: %s. Using single rank timing.", e)
2025-09-07 14:59:17 +08:00
return rank_name_to_time
def _get_global_min_max_time(self, names, reset, barrier, normalizer):
rank_name_to_time = self._get_elapsed_time_all_ranks(names, reset, barrier)
name_to_min_max_time = {}
for i, name in enumerate(names):
rank_to_time = rank_name_to_time[:, i]
rank_to_time = rank_to_time[rank_to_time > 0.0]
if rank_to_time.numel() > 0:
name_to_min_max_time[name] = (
rank_to_time.min().item() / normalizer,
rank_to_time.max().item() / normalizer,
)
return name_to_min_max_time
2025-09-11 13:18:33 +08:00
def _get_global_min_max_time_string(
self, names, reset, barrier, normalizer, max_only
):
name_to_min_max_time = self._get_global_min_max_time(
names, reset, barrier, normalizer
)
2025-09-07 14:59:17 +08:00
if not name_to_min_max_time:
return None
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
world_size = _get_world_size()
if world_size == 1:
2025-09-11 13:18:33 +08:00
output_string = "time (ms):"
2025-09-07 14:59:17 +08:00
for name in name_to_min_max_time:
2026-06-15 11:40:00 +08:00
_, max_time = name_to_min_max_time[name]
2025-09-11 13:18:33 +08:00
output_string += "\n {}: {:.2f}".format(
(name + " ").ljust(48, "."), max_time
)
2025-09-07 14:59:17 +08:00
else:
if max_only:
2025-09-11 13:18:33 +08:00
output_string = "max time across ranks (ms):"
2025-09-07 14:59:17 +08:00
else:
2025-09-11 13:18:33 +08:00
output_string = "(min, max) time across ranks (ms):"
2025-09-07 14:59:17 +08:00
for name in name_to_min_max_time:
min_time, max_time = name_to_min_max_time[name]
if max_only:
2025-09-11 13:18:33 +08:00
output_string += "\n {}: {:.2f}".format(
(name + " ").ljust(48, "."), max_time
)
2025-09-07 14:59:17 +08:00
else:
2025-09-11 13:18:33 +08:00
output_string += "\n {}: ({:.2f}, {:.2f})".format(
(name + " ").ljust(48, "."), min_time, max_time
2025-09-07 14:59:17 +08:00
)
return output_string
def _get_all_ranks_time_string(self, names, reset, barrier, normalizer):
rank_name_to_time = self._get_elapsed_time_all_ranks(names, reset, barrier)
world_size = _get_world_size()
2025-09-11 13:18:33 +08:00
output_string = "times across ranks (ms):"
2025-09-07 14:59:17 +08:00
no_reported_timing = True
for i, name in enumerate(names):
not_yet_found = True
for rank in range(world_size):
if rank_name_to_time[rank, i] > 0:
no_reported_timing = False
if not_yet_found:
not_yet_found = False
2025-09-11 13:18:33 +08:00
output_string += "\n {}:".format(name)
2025-09-07 14:59:17 +08:00
if world_size == 1:
2025-09-11 13:18:33 +08:00
output_string += "\n {:.2f}".format(
2025-09-07 14:59:17 +08:00
rank_name_to_time[rank, i] / normalizer
)
else:
2025-09-11 13:18:33 +08:00
output_string += "\n rank {:2d}: {:.2f}".format(
2025-09-07 14:59:17 +08:00
rank, rank_name_to_time[rank, i] / normalizer
)
if no_reported_timing:
return None
return output_string
def get_all_timers_string(
self,
names: List[str] = None,
normalizer: float = 1.0,
reset: bool = True,
barrier: bool = False,
):
2026-06-15 11:40:00 +08:00
"""Return a formatted timing string for the given timer names.
2025-09-07 14:59:17 +08:00
Args:
2026-06-15 11:40:00 +08:00
names: Timers to include; defaults to all registered timers.
normalizer: Divide raw seconds by this value (e.g. 1000 for ms output).
reset: Reset each timer after reading its elapsed time.
barrier: Synchronize across ranks before gathering times.
2025-09-07 14:59:17 +08:00
"""
2026-06-15 11:40:00 +08:00
if names is None:
2025-09-07 14:59:17 +08:00
names = list(self._timers.keys())
assert normalizer > 0.0
2025-09-11 13:18:33 +08:00
if self._log_option in ["max", "minmax"]:
2026-06-15 11:40:00 +08:00
max_only = self._log_option == "max"
2025-09-07 14:59:17 +08:00
output_string = self._get_global_min_max_time_string(
names, reset, barrier, normalizer / 1000.0, max_only
)
2025-09-11 13:18:33 +08:00
elif self._log_option == "all":
2025-09-07 14:59:17 +08:00
output_string = self._get_all_ranks_time_string(
names, reset, barrier, normalizer / 1000.0
)
else:
2025-09-11 13:18:33 +08:00
raise Exception("unknown timing log option {}".format(self._log_option))
2025-09-07 14:59:17 +08:00
return output_string
def log(
self,
names: List[str],
rank: int = None,
normalizer: float = 1.0,
reset: bool = True,
barrier: bool = False,
):
2026-06-15 11:40:00 +08:00
"""Print timing results for the given names to stdout on one rank.
2025-09-07 14:59:17 +08:00
Args:
2026-06-15 11:40:00 +08:00
names: Timer names to log.
rank: Rank that prints; defaults to the last rank (world_size - 1).
normalizer: Divide raw seconds by this value before printing.
reset: Reset each timer after reading.
barrier: Synchronize across ranks first.
2025-09-07 14:59:17 +08:00
"""
output_string = self.get_all_timers_string(names, normalizer, reset, barrier)
world_size = _get_world_size()
current_rank = _get_rank()
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
if rank is None:
rank = world_size - 1
if rank == current_rank and output_string is not None:
2026-06-15 11:40:00 +08:00
logger.info("%s", output_string)
2025-09-07 14:59:17 +08:00
def write(
self,
names: List[str],
writer,
iteration: int,
normalizer: float = 1.0,
reset: bool = True,
barrier: bool = False,
):
2026-06-15 11:40:00 +08:00
"""Write per-timer max times as TensorBoard scalars.
2025-09-07 14:59:17 +08:00
Args:
2026-06-15 11:40:00 +08:00
names: Timer names to write.
writer: TensorBoard SummaryWriter instance.
iteration: Global step value for the scalar.
normalizer: Divide raw seconds by this value.
reset: Reset each timer after reading.
barrier: Synchronize across ranks first.
2025-09-07 14:59:17 +08:00
"""
assert normalizer > 0.0
2025-09-11 13:18:33 +08:00
name_to_min_max_time = self._get_global_min_max_time(
names, reset, barrier, normalizer
)
2025-09-07 14:59:17 +08:00
if writer is not None:
for name in name_to_min_max_time:
_, max_time = name_to_min_max_time[name]
2025-09-11 13:18:33 +08:00
writer.add_scalar(name + "-time", max_time, iteration)