Initial import: grid-bot — grid trading bot for BTC-USDT on Cifra Markets
This commit is contained in:
@@ -0,0 +1,60 @@
|
||||
r"""
|
||||
PyTorch Profiler is a tool that allows the collection of performance metrics during training and inference.
|
||||
Profiler's context manager API can be used to better understand what model operators are the most expensive,
|
||||
examine their input shapes and stack traces, study device kernel activity and visualize the execution trace.
|
||||
|
||||
.. note::
|
||||
An earlier version of the API in :mod:`torch.autograd` module is considered legacy and will be deprecated.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
from typing_extensions import TypeVarTuple, Unpack
|
||||
|
||||
from torch._C._autograd import _supported_activities, DeviceType, kineto_available
|
||||
from torch._C._profiler import _ExperimentalConfig, ProfilerActivity, RecordScope
|
||||
from torch._environment import is_fbcode
|
||||
from torch.autograd.profiler import KinetoStepTracker, record_function
|
||||
from torch.optim.optimizer import Optimizer, register_optimizer_step_post_hook
|
||||
|
||||
from .profiler import (
|
||||
_KinetoProfile,
|
||||
ExecutionTraceObserver,
|
||||
profile,
|
||||
ProfilerAction,
|
||||
schedule,
|
||||
supported_activities,
|
||||
tensorboard_trace_handler,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"profile",
|
||||
"schedule",
|
||||
"supported_activities",
|
||||
"tensorboard_trace_handler",
|
||||
"ProfilerAction",
|
||||
"ProfilerActivity",
|
||||
"kineto_available",
|
||||
"DeviceType",
|
||||
"record_function",
|
||||
"ExecutionTraceObserver",
|
||||
]
|
||||
|
||||
from . import itt
|
||||
|
||||
|
||||
_Ts = TypeVarTuple("_Ts")
|
||||
|
||||
|
||||
def _optimizer_post_hook(
|
||||
optimizer: Optimizer, args: tuple[Unpack[_Ts]], kwargs: dict[str, Any]
|
||||
) -> None:
|
||||
KinetoStepTracker.increment_step("Optimizer")
|
||||
|
||||
|
||||
if os.environ.get("KINETO_USE_DAEMON", "") or (
|
||||
is_fbcode() and os.environ.get("KINETO_FORCE_OPTIMIZER_HOOK", "")
|
||||
):
|
||||
_ = register_optimizer_step_post_hook(_optimizer_post_hook)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,681 @@
|
||||
# mypy: allow-untyped-defs
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
|
||||
import torch
|
||||
import torch.utils.benchmark as benchmark
|
||||
from torch._C._profiler import (
|
||||
_EventType,
|
||||
_ExtraFields_PyCall,
|
||||
_ExtraFields_PyCCall,
|
||||
_ExtraFields_TorchOp,
|
||||
_ProfilerEvent,
|
||||
)
|
||||
from torch.profiler import profile
|
||||
from torch.profiler._utils import index_of_first_match, traverse_bfs, traverse_dfs
|
||||
|
||||
|
||||
class Pattern:
|
||||
"""
|
||||
Base class for all patterns, subclass this class and implement match()
|
||||
to define custom patterns.
|
||||
|
||||
In subclass, define description and skip property.
|
||||
"""
|
||||
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
self.prof = prof
|
||||
self.should_benchmark = should_benchmark
|
||||
self.name = "Please specify a name for pattern"
|
||||
self.description = "Please specify a description for pattern"
|
||||
self.url = ""
|
||||
if prof.profiler is None or prof.profiler.kineto_results is None:
|
||||
raise AssertionError("profiler and kineto_results must not be None")
|
||||
self.event_tree = prof.profiler.kineto_results.experimental_event_tree()
|
||||
self.tid_root: dict[int, list[_ProfilerEvent]] = {}
|
||||
for event in self.event_tree:
|
||||
self.tid_root.setdefault(event.start_tid, []).append(event)
|
||||
|
||||
@property
|
||||
def skip(self) -> bool:
|
||||
return False
|
||||
|
||||
def report(self, event: _ProfilerEvent):
|
||||
msg = (
|
||||
f"{self.description}\n[Source Code Location] {source_code_location(event)}"
|
||||
)
|
||||
return msg
|
||||
|
||||
def eventTreeTraversal(self):
|
||||
"""
|
||||
Traverse the event tree and yield all events.
|
||||
Override this method in subclass to customize the traversal.
|
||||
"""
|
||||
yield from traverse_dfs(self.event_tree)
|
||||
|
||||
def summary(self, events: list[_ProfilerEvent]):
|
||||
default_summary = f"{self.name}: {len(events)} events matched."
|
||||
if self.should_benchmark:
|
||||
# If benchmark summary is not empty, use it.
|
||||
return (
|
||||
self.benchmark_summary(events)
|
||||
if hasattr(self, "benchmark") # type: ignore[attr-defined]
|
||||
else default_summary
|
||||
)
|
||||
return default_summary
|
||||
|
||||
def benchmark_summary(self, events: list[_ProfilerEvent]) -> str:
|
||||
def format_time(time_ns: int) -> str:
|
||||
unit_lst = ["ns", "us", "ms"]
|
||||
for unit in unit_lst:
|
||||
if time_ns < 1000:
|
||||
return f"{time_ns:.2f} {unit}"
|
||||
time_ns //= 1000
|
||||
return f"{time_ns:.2f} s"
|
||||
|
||||
if not hasattr(self, "benchmark"):
|
||||
raise AssertionError("Please implement benchmark()")
|
||||
shapes_factor_map = self.benchmark(events) # type: ignore[attr-defined]
|
||||
original_time = sum(event.duration_time_ns for event in events)
|
||||
new_time = sum(
|
||||
shapes_factor_map[input_shapes(event)] * event.duration_time_ns
|
||||
for event in events
|
||||
)
|
||||
return (
|
||||
f"{self.name}: {len(events)} events matched. "
|
||||
f"Total Estimated Speedup: {format_time(original_time - new_time)} ({round(original_time / new_time, 2)}X)"
|
||||
)
|
||||
|
||||
def match(self, event: _ProfilerEvent):
|
||||
"""
|
||||
Return True if the event matches the pattern.
|
||||
This method should be overridden in subclass.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def matched_events(self):
|
||||
if self.skip:
|
||||
return []
|
||||
matched_events = [
|
||||
event for event in self.eventTreeTraversal() if self.match(event)
|
||||
]
|
||||
return matched_events
|
||||
|
||||
def root_of(self, event: _ProfilerEvent):
|
||||
while event.parent:
|
||||
event = event.parent
|
||||
return event
|
||||
|
||||
def siblings_of(self, event: _ProfilerEvent):
|
||||
if event.parent:
|
||||
children = event.parent.children
|
||||
else:
|
||||
children = self.tid_root[event.start_tid]
|
||||
index = children.index(event)
|
||||
return children[:index], children[index + 1 :]
|
||||
|
||||
def next_of(self, event: _ProfilerEvent):
|
||||
_, next_events = self.siblings_of(event)
|
||||
return next_events[0] if next_events else None
|
||||
|
||||
def prev_of(self, event: _ProfilerEvent):
|
||||
prev_events, _ = self.siblings_of(event)
|
||||
return prev_events[-1] if prev_events else None
|
||||
|
||||
def go_up_until(self, event: _ProfilerEvent, predicate):
|
||||
if not event:
|
||||
return None
|
||||
while event.parent and not predicate(event):
|
||||
event = event.parent
|
||||
return event
|
||||
|
||||
|
||||
# Patterns
|
||||
|
||||
|
||||
class NamePattern(Pattern):
|
||||
def __init__(
|
||||
self, prof: profile, name: str, should_benchmark: bool = False
|
||||
) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.description = f"Matched Name Event: {name}"
|
||||
self.name = name
|
||||
|
||||
def match(self, event: _ProfilerEvent):
|
||||
return re.search(self.name, event.name) is not None
|
||||
|
||||
|
||||
class ExtraCUDACopyPattern(Pattern):
|
||||
"""
|
||||
This pattern identifies if we creates a constant tensor on CPU and immediately moves it to GPU.
|
||||
example: torch.zeros((100, 100)).to("cuda")
|
||||
|
||||
Pattern:
|
||||
built-in method |built-in method
|
||||
... | aten::to
|
||||
aten::fill_/aten::zero_ | aten::_to_copy
|
||||
|
||||
Algorithm:
|
||||
We start at node aten::to, go parent events' previous events,
|
||||
and check if we have a aten::fill_/aten::zero_ as we keep going down the tree.
|
||||
We always select the last child in the children list when we go down the tree.
|
||||
If at any step we failed, it is not a match.
|
||||
"""
|
||||
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "Extra CUDA Copy Pattern"
|
||||
self.description = "Filled a CPU tensor and immediately moved it to GPU. Please initialize it on GPU."
|
||||
self.url = "https://pytorch.org/tutorials/recipes/recipes/tuning_guide.html#create-tensors-directly-on-the-target-device"
|
||||
self.init_ops = {
|
||||
"aten::fill_",
|
||||
"aten::zero_",
|
||||
"aten::normal_",
|
||||
"aten::uniform_",
|
||||
}
|
||||
|
||||
@property
|
||||
def skip(self) -> bool:
|
||||
return not self.prof.with_stack or not self.prof.record_shapes
|
||||
|
||||
def match(self, event):
|
||||
# TODO: We should also check tensor identities
|
||||
if event.name != "aten::to":
|
||||
return False
|
||||
to_event = event
|
||||
if not event.children:
|
||||
return False
|
||||
event = event.children[-1]
|
||||
if event.name != "aten::_to_copy":
|
||||
return False
|
||||
if not event.children:
|
||||
return False
|
||||
event = event.children[-1]
|
||||
if event.name != "aten::copy_":
|
||||
return False
|
||||
# aten::copy_ should have the first 2 args dtype the same
|
||||
dtypes = input_dtypes(event)
|
||||
if len(dtypes) < 2:
|
||||
return False
|
||||
if dtypes[0] is None or dtypes[0] != dtypes[1]:
|
||||
return False
|
||||
event = to_event
|
||||
# Up one level
|
||||
event = event.parent
|
||||
if event is None:
|
||||
return False
|
||||
# Check if we have a aten::fill_ in previous leaf
|
||||
event = self.prev_of(event)
|
||||
if event is None:
|
||||
return False
|
||||
while event.children:
|
||||
event = event.children[-1]
|
||||
# aten::zero_ is a special optimization case where fill_ is not called
|
||||
if event.name in self.init_ops:
|
||||
return True
|
||||
return event.name in self.init_ops
|
||||
# TODO: Check if tensor is reused
|
||||
|
||||
def benchmark(self, events: list[_ProfilerEvent]):
|
||||
shapes_factor_map = {input_shapes(event): 0.0 for event in events}
|
||||
for shape in shapes_factor_map:
|
||||
size = shape[0]
|
||||
to_timer = benchmark.Timer(
|
||||
stmt='torch.ones(size).to("cuda")', globals={"size": size}
|
||||
)
|
||||
de_timer = benchmark.Timer(
|
||||
stmt='torch.ones(size, device="cuda")', globals={"size": size}
|
||||
)
|
||||
to_time = to_timer.timeit(10).mean
|
||||
de_time = de_timer.timeit(10).mean
|
||||
shapes_factor_map[shape] = de_time / to_time
|
||||
return shapes_factor_map
|
||||
|
||||
|
||||
class ForLoopIndexingPattern(Pattern):
|
||||
"""
|
||||
This pattern identifies if we use a for loop to index a tensor that
|
||||
can be vectorized.
|
||||
example:
|
||||
tensor = torch.empty((100, 100))
|
||||
for i in range(100):
|
||||
tensor[i] = i
|
||||
|
||||
Pattern:
|
||||
aten::select | ... | aten::select | ... (Repeat)
|
||||
|
||||
Algorithm:
|
||||
We start at node aten::select, and we check if we can find this alternating patterns.
|
||||
We also keep a dictionary to avoid duplicate match in the for loop.
|
||||
"""
|
||||
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "For Loop Indexing Pattern"
|
||||
self.description = "For loop indexing detected. Vectorization recommended."
|
||||
self.visited: set[int] = set()
|
||||
|
||||
def eventTreeTraversal(self):
|
||||
"""
|
||||
We need to use BFS traversal order to avoid duplicate match.
|
||||
"""
|
||||
yield from traverse_bfs(self.event_tree)
|
||||
|
||||
def match(self, event: _ProfilerEvent):
|
||||
if event.name != "aten::select":
|
||||
return False
|
||||
if event.id in self.visited:
|
||||
return False
|
||||
repeat_count = 1
|
||||
_, next = self.siblings_of(event)
|
||||
if len(next) <= 1:
|
||||
return False
|
||||
|
||||
# Custom event list matching
|
||||
def same_ops(list1, list2) -> bool:
|
||||
if len(list1) != len(list2):
|
||||
return False
|
||||
for op1, op2 in zip(list1, list2, strict=True):
|
||||
if op1.name != op2.name:
|
||||
return False
|
||||
return True
|
||||
|
||||
# Record the ops between two aten::select
|
||||
next_select_idx = index_of_first_match(next, lambda e: e.name == "aten::select")
|
||||
if next_select_idx is None:
|
||||
return False
|
||||
indexing_ops = [event] + next[:next_select_idx]
|
||||
next = next[len(indexing_ops) - 1 :]
|
||||
for i in range(0, len(next), len(indexing_ops)):
|
||||
if same_ops(indexing_ops, next[i : i + len(indexing_ops)]):
|
||||
repeat_count += 1
|
||||
self.visited.add(next[i].id)
|
||||
else:
|
||||
break
|
||||
return repeat_count >= 10
|
||||
|
||||
|
||||
class FP32MatMulPattern(Pattern):
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "FP32 MatMul Pattern"
|
||||
self.description = (
|
||||
"You are currently using GPU that supports TF32. "
|
||||
"Please enable TF32 by setting 'torch.backends.cuda.matmul.allow_tf32 = True'"
|
||||
)
|
||||
self.url = "https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
|
||||
|
||||
@property
|
||||
def skip(self):
|
||||
if torch.version.hip is not None:
|
||||
has_tf32 = False
|
||||
else:
|
||||
# Anything less than sm_80 is not Ampere which doesn't support TF32
|
||||
has_tf32 = all(
|
||||
int(re.sub("sm_|compute_", "", arch)) >= 80
|
||||
for arch in torch.cuda.get_arch_list()
|
||||
)
|
||||
return has_tf32 is False or super().skip or not self.prof.record_shapes
|
||||
|
||||
def match(self, event: _ProfilerEvent) -> bool:
|
||||
# If we saw this pattern once, we don't need to match it again
|
||||
if event.tag != _EventType.TorchOp:
|
||||
return False
|
||||
if not isinstance(event.extra_fields, _ExtraFields_TorchOp):
|
||||
raise AssertionError(
|
||||
f"expected _ExtraFields_TorchOp, got {type(event.extra_fields).__name__}"
|
||||
)
|
||||
if event.name == "aten::mm":
|
||||
if event.extra_fields.allow_tf32_cublas is False:
|
||||
return True
|
||||
return False
|
||||
|
||||
def report(self, event: _ProfilerEvent):
|
||||
return self.description
|
||||
|
||||
def benchmark(self, events: list[_ProfilerEvent]):
|
||||
shapes_factor_map = {input_shapes(event): 0.0 for event in events}
|
||||
for shape in shapes_factor_map:
|
||||
matrixA = torch.randn(shape[0], device="cuda", dtype=torch.float32)
|
||||
matrixB = torch.randn(shape[1], device="cuda", dtype=torch.float32)
|
||||
fp32_timer = benchmark.Timer(
|
||||
stmt="torch.mm(matrixA, matrixB)",
|
||||
globals={"matrixA": matrixA, "matrixB": matrixB},
|
||||
)
|
||||
tf32_timer = benchmark.Timer(
|
||||
stmt="torch.mm(matrixA, matrixB)",
|
||||
setup="torch.backends.cuda.matmul.allow_tf32 = True",
|
||||
globals={"matrixA": matrixA, "matrixB": matrixB},
|
||||
)
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
fp32_time = fp32_timer.timeit(10).mean
|
||||
tf32_time = tf32_timer.timeit(10).mean
|
||||
shapes_factor_map[shape] = tf32_time / fp32_time
|
||||
return shapes_factor_map
|
||||
|
||||
|
||||
class OptimizerSingleTensorPattern(Pattern):
|
||||
"""
|
||||
This pattern identifies if we are using the single-tensor version of an optimizer.
|
||||
example:
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
|
||||
By adding foreach=True to enable multi-tensor optimizer, we can gain speedup when
|
||||
the kernels are relatively small.
|
||||
|
||||
Pattern:
|
||||
XXXXX: _single_tenser_<OPTIMIZER_NAME>
|
||||
|
||||
Algorithm:
|
||||
String match
|
||||
"""
|
||||
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "Optimizer Single Tensor Pattern"
|
||||
self.optimizers_with_foreach = ["adam", "sgd", "adamw"]
|
||||
self.description = (
|
||||
"Detected optimizer running with single tensor implementation. "
|
||||
"Please enable multi tensor implementation by passing 'foreach=True' into optimizer."
|
||||
)
|
||||
self.url = ""
|
||||
|
||||
def match(self, event: _ProfilerEvent) -> bool:
|
||||
for optimizer in self.optimizers_with_foreach:
|
||||
if event.name.endswith(f"_single_tensor_{optimizer}"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class SynchronizedDataLoaderPattern(Pattern):
|
||||
"""
|
||||
This pattern identifies if we are using num_workers=0 in DataLoader.
|
||||
example:
|
||||
torch.utils.data.DataLoader(dataset, batch_size=batch_size)
|
||||
Add num_workers=N to the arguments. N depends on system configuration.
|
||||
|
||||
Pattern:
|
||||
dataloader.py(...): __iter__
|
||||
dataloader.py(...): _get_iterator
|
||||
NOT dataloader.py(...): check_worker_number_rationality
|
||||
|
||||
Algorithm:
|
||||
If we don't see check_worker_number_rationality call in the dataloader __iter__,
|
||||
It is not an asynchronous dataloader.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "Synchronized DataLoader Pattern"
|
||||
self.description = (
|
||||
"Detected DataLoader running with synchronized implementation. "
|
||||
"Please enable asynchronous dataloading by setting num_workers > 0 when initializing DataLoader."
|
||||
)
|
||||
self.url = (
|
||||
"https://pytorch.org/tutorials/recipes/recipes/tuning_guide.html"
|
||||
"#enable-async-data-loading-and-augmentation"
|
||||
)
|
||||
|
||||
def match(self, event: _ProfilerEvent) -> bool:
|
||||
def is_dataloader_function(name: str, function_name: str):
|
||||
return name.startswith(
|
||||
os.path.join("torch", "utils", "data", "dataloader.py")
|
||||
) and name.endswith(function_name)
|
||||
|
||||
# TODO: fixme! Due to lifetime issues of the function name, this field might
|
||||
# actually point to an already freed string when the even is a PyCall.
|
||||
# Just silently skip this to unblock testing.
|
||||
try:
|
||||
event.name
|
||||
except UnicodeDecodeError:
|
||||
return False
|
||||
|
||||
if not is_dataloader_function(event.name, "__iter__"):
|
||||
return False
|
||||
if not event.children:
|
||||
return False
|
||||
event = event.children[0]
|
||||
if not is_dataloader_function(event.name, "_get_iterator"):
|
||||
return False
|
||||
if not event.children:
|
||||
return False
|
||||
event = event.children[0]
|
||||
return not is_dataloader_function(event.name, "check_worker_number_rationality")
|
||||
# TODO: We should also check if the loader is bottleneck.
|
||||
|
||||
|
||||
class GradNotSetToNonePattern(Pattern):
|
||||
"""
|
||||
This pattern identifies if we are not setting grad to None in zero_grad.
|
||||
example:
|
||||
optimizer.zero_grad()
|
||||
By setting set_to_none=True, we can gain speedup
|
||||
|
||||
Pattern:
|
||||
XXXXX: _zero_grad
|
||||
NOT aten::zeros
|
||||
aten::zero_
|
||||
|
||||
aten::zero_ is called on each parameter in the model.
|
||||
We also want to make sure it is not called by aten::zeros.
|
||||
|
||||
Algorithm:
|
||||
String match
|
||||
"""
|
||||
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "Gradient Set To Zero Instead of None Pattern"
|
||||
self.description = (
|
||||
"Detected gradient set to zero instead of None. "
|
||||
"Please add 'set_to_none=True' when calling zero_grad()."
|
||||
)
|
||||
self.url = (
|
||||
"https://pytorch.org/tutorials/recipes/recipes/tuning_guide.html"
|
||||
"#disable-gradient-calculation-for-validation-or-inference"
|
||||
)
|
||||
|
||||
def match(self, event: _ProfilerEvent) -> bool:
|
||||
if not event.name.endswith(": zero_grad"):
|
||||
return False
|
||||
if not event.children:
|
||||
return False
|
||||
|
||||
for sub_event in traverse_dfs(event.children):
|
||||
if (
|
||||
sub_event.name == "aten::zero_"
|
||||
and sub_event.parent.name != "aten::zeros"
|
||||
):
|
||||
return True
|
||||
# TODO: We should also check if the optimizer's numerical behavior will change.
|
||||
return False
|
||||
|
||||
|
||||
class Conv2dBiasFollowedByBatchNorm2dPattern(Pattern):
|
||||
"""
|
||||
This pattern identifies if we are enabling bias in Conv2d which is followed by BatchNorm2d.
|
||||
Bias doesn't do anything when followed by batchnorm.
|
||||
Pattern:
|
||||
nn.Module: Conv2d | nn.Module: BatchNorm2d
|
||||
...
|
||||
aten::conv2d AND dtype of third argument is not null
|
||||
The third argument is the bias
|
||||
Algorithm:
|
||||
String match
|
||||
"""
|
||||
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "Enabling Bias in Conv2d Followed By BatchNorm Pattern"
|
||||
self.description = "Detected bias enabled in Conv2d that is followed by BatchNorm2d. Please set 'bias=False' in Conv2d."
|
||||
self.url = (
|
||||
"https://pytorch.org/tutorials/recipes/recipes/tuning_guide.html"
|
||||
"#disable-bias-for-convolutions-directly-followed-by-a-batch-norm"
|
||||
)
|
||||
|
||||
@property
|
||||
def skip(self):
|
||||
return self.prof.record_shapes is False or super().skip
|
||||
|
||||
def match(self, event: _ProfilerEvent):
|
||||
if event.name != "aten::conv2d":
|
||||
return False
|
||||
if len(input_dtypes(event)) < 3 or input_dtypes(event)[2] is None:
|
||||
return False
|
||||
# This means bias=True
|
||||
event = self.go_up_until(
|
||||
event, lambda e: e.name.startswith("nn.Module: Conv2d")
|
||||
)
|
||||
if not event:
|
||||
return False
|
||||
event = self.next_of(event)
|
||||
if not event:
|
||||
return False
|
||||
return event.name.startswith("nn.Module: BatchNorm2d")
|
||||
|
||||
|
||||
class MatMulDimInFP16Pattern(Pattern):
|
||||
def __init__(self, prof: profile, should_benchmark: bool = False) -> None:
|
||||
super().__init__(prof, should_benchmark)
|
||||
self.name = "Matrix Multiplication Dimension Not Aligned Pattern"
|
||||
self.description = "Detected matmul with dimension not aligned. Please use matmul with aligned dimension."
|
||||
self.url = "https://pytorch.org/tutorials/recipes/recipes/tuning_guide.html#use-mixed-precision-and-amp"
|
||||
|
||||
@property
|
||||
def skip(self) -> bool:
|
||||
return not self.prof.with_stack or not self.prof.record_shapes
|
||||
|
||||
def match(self, event: _ProfilerEvent) -> bool:
|
||||
def mutiple_of(shapes, multiple):
|
||||
return all(dim % multiple == 0 for shape in shapes for dim in shape[-2:])
|
||||
|
||||
if event.name not in ("aten::mm", "aten::bmm", "aten::addmm"):
|
||||
return False
|
||||
if not input_dtypes(event):
|
||||
return False
|
||||
arg_dtype = input_dtypes(event)[0]
|
||||
if arg_dtype in (torch.bfloat16, torch.half) and not mutiple_of(
|
||||
input_shapes(event), 8
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
def benchmark(self, events: list[_ProfilerEvent]):
|
||||
def closest_multiple(shapes, multiple):
|
||||
return [multiple * math.ceil(shape / multiple) for shape in shapes]
|
||||
|
||||
shapes_factor_map = {input_shapes(event): 0.0 for event in events}
|
||||
for shape in shapes_factor_map:
|
||||
matrixA = torch.randn(shape[0], device="cuda", dtype=torch.float16)
|
||||
matrixB = torch.randn(shape[1], device="cuda", dtype=torch.float16)
|
||||
not_aligned_dim_timer = benchmark.Timer(
|
||||
stmt="torch.mm(matrixA, matrixB)",
|
||||
globals={"matrixA": matrixA, "matrixB": matrixB},
|
||||
)
|
||||
matrixA = torch.randn(
|
||||
closest_multiple(shape[0], 8), device="cuda", dtype=torch.float16
|
||||
)
|
||||
matrixB = torch.randn(
|
||||
closest_multiple(shape[1], 8), device="cuda", dtype=torch.float16
|
||||
)
|
||||
aligned_dim_timer = benchmark.Timer(
|
||||
stmt="torch.mm(matrixA, matrixB)",
|
||||
globals={"matrixA": matrixA, "matrixB": matrixB},
|
||||
)
|
||||
not_aligned_dim_time = not_aligned_dim_timer.timeit(10).mean
|
||||
aligned_dim_time = aligned_dim_timer.timeit(10).mean
|
||||
shapes_factor_map[shape] = aligned_dim_time / not_aligned_dim_time
|
||||
return shapes_factor_map
|
||||
|
||||
|
||||
def source_code_location(event: _ProfilerEvent | None) -> str:
|
||||
while event:
|
||||
if event.tag == _EventType.PyCall or event.tag == _EventType.PyCCall:
|
||||
if not isinstance(
|
||||
event.extra_fields, (_ExtraFields_PyCall, _ExtraFields_PyCCall)
|
||||
):
|
||||
raise AssertionError(
|
||||
f"expected _ExtraFields_PyCall or _ExtraFields_PyCCall, "
|
||||
f"got {type(event.extra_fields).__name__}"
|
||||
)
|
||||
if not event.extra_fields.caller.file_name.startswith("torch" + os.sep):
|
||||
return f"{event.extra_fields.caller.file_name}:{event.extra_fields.caller.line_number}"
|
||||
event = event.parent
|
||||
return "No source code location found"
|
||||
|
||||
|
||||
def input_shapes(event: _ProfilerEvent):
|
||||
if not isinstance(event.extra_fields, _ExtraFields_TorchOp):
|
||||
raise AssertionError(
|
||||
f"expected _ExtraFields_TorchOp, got {type(event.extra_fields).__name__}"
|
||||
)
|
||||
return tuple(tuple(getattr(i, "sizes", ())) for i in event.extra_fields.inputs)
|
||||
|
||||
|
||||
def input_dtypes(event: _ProfilerEvent):
|
||||
if not isinstance(event.extra_fields, _ExtraFields_TorchOp):
|
||||
raise AssertionError(
|
||||
f"expected _ExtraFields_TorchOp, got {type(event.extra_fields).__name__}"
|
||||
)
|
||||
return tuple(getattr(i, "dtype", None) for i in event.extra_fields.inputs)
|
||||
|
||||
|
||||
def report_all_anti_patterns(
|
||||
prof,
|
||||
should_benchmark: bool = False,
|
||||
print_enable: bool = True,
|
||||
json_report_dir: str | None = None,
|
||||
) -> None:
|
||||
report_dict: dict = {}
|
||||
anti_patterns = [
|
||||
ExtraCUDACopyPattern(prof, should_benchmark),
|
||||
# ForLoopIndexingPattern(prof, should_benchmark),
|
||||
FP32MatMulPattern(prof, should_benchmark),
|
||||
OptimizerSingleTensorPattern(prof, should_benchmark),
|
||||
SynchronizedDataLoaderPattern(prof, should_benchmark),
|
||||
GradNotSetToNonePattern(prof, should_benchmark),
|
||||
Conv2dBiasFollowedByBatchNorm2dPattern(prof, should_benchmark),
|
||||
MatMulDimInFP16Pattern(prof, should_benchmark),
|
||||
]
|
||||
reported = set()
|
||||
summaries = []
|
||||
message_list = [f"{'-' * 40}TorchTidy Report{'-' * 40}"]
|
||||
message_list.append("Matched Events:")
|
||||
|
||||
for anti_pattern in anti_patterns:
|
||||
matched_events = anti_pattern.matched_events()
|
||||
if not matched_events:
|
||||
continue
|
||||
summaries.append(anti_pattern.summary(matched_events))
|
||||
for event in matched_events:
|
||||
report_msg = anti_pattern.report(event)
|
||||
if report_msg not in reported:
|
||||
message_list.append(report_msg)
|
||||
reported.add(report_msg)
|
||||
src_location, line_no = source_code_location(event).split(":")
|
||||
report_dict.setdefault(src_location, []).append(
|
||||
{
|
||||
"line_number": int(line_no),
|
||||
"name": anti_pattern.name,
|
||||
"url": anti_pattern.url,
|
||||
"message": anti_pattern.description,
|
||||
}
|
||||
)
|
||||
|
||||
if json_report_dir is not None:
|
||||
json_report_path = os.path.join(json_report_dir, "torchtidy_report.json")
|
||||
if os.path.exists(json_report_path):
|
||||
with open(json_report_path) as f:
|
||||
exisiting_report = json.load(f)
|
||||
exisiting_report.update(report_dict)
|
||||
report_dict = exisiting_report
|
||||
with open(json_report_path, "w") as f:
|
||||
json.dump(report_dict, f, indent=4)
|
||||
|
||||
message_list.append("Summary:")
|
||||
message_list += summaries
|
||||
message_list.append(f"{'-' * 40}TorchTidy Report{'-' * 40}")
|
||||
if print_enable:
|
||||
print("\n".join(message_list))
|
||||
@@ -0,0 +1,577 @@
|
||||
# mypy: allow-untyped-defs
|
||||
import functools
|
||||
import operator
|
||||
import re
|
||||
from collections import deque
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal, TYPE_CHECKING
|
||||
|
||||
from torch.autograd.profiler import profile
|
||||
from torch.profiler import DeviceType
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from torch.autograd import _KinetoEvent
|
||||
|
||||
|
||||
def _traverse(tree, next_fn, children_fn=lambda x: x.children, reverse: bool = False):
|
||||
order = reversed if reverse else lambda x: x
|
||||
remaining = deque(order(tree))
|
||||
while remaining:
|
||||
curr_event = next_fn(remaining)
|
||||
yield curr_event
|
||||
for child_event in order(children_fn(curr_event)):
|
||||
remaining.append(child_event)
|
||||
|
||||
|
||||
traverse_dfs = functools.partial(_traverse, next_fn=lambda x: x.pop(), reverse=True)
|
||||
traverse_bfs = functools.partial(
|
||||
_traverse, next_fn=lambda x: x.popleft(), reverse=False
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class EventMetrics:
|
||||
duration_time_ns: int = 0
|
||||
self_time_ns: int = 0
|
||||
idle_time_ns: int = 0
|
||||
queue_depth: int = 0
|
||||
|
||||
@property
|
||||
def fraction_idle_time(self):
|
||||
if self.duration_time_ns == 0:
|
||||
return 0.0
|
||||
return self.idle_time_ns / self.duration_time_ns
|
||||
|
||||
|
||||
@dataclass
|
||||
class Interval:
|
||||
start: int
|
||||
end: int
|
||||
queue_depth: int = 0
|
||||
|
||||
|
||||
class EventKey:
|
||||
def __init__(self, event) -> None:
|
||||
self.event = event
|
||||
|
||||
def __hash__(self):
|
||||
return hash(self.event.id)
|
||||
|
||||
def __eq__(self, other):
|
||||
return self.event.id == other.event.id
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.event.name}"
|
||||
|
||||
def intervals_overlap(self, intervals: list[Interval]):
|
||||
overlap_time = 0
|
||||
intervals = sorted(intervals, key=lambda x: x.start)
|
||||
|
||||
if intervals:
|
||||
overlap_start = max(self.event.start_time_ns, intervals[0].start)
|
||||
overlap_end = min(self.event.end_time_ns, intervals[0].end)
|
||||
|
||||
if overlap_start < overlap_end:
|
||||
overlap_time += overlap_end - overlap_start
|
||||
|
||||
i, j = 0, 1
|
||||
while j < len(intervals):
|
||||
prev_interval = intervals[i]
|
||||
curr_interval = intervals[j]
|
||||
j += 1
|
||||
if prev_interval.end > curr_interval.start:
|
||||
# Completely subsumed by previous interval
|
||||
if prev_interval.end > curr_interval.end:
|
||||
j += 1
|
||||
continue
|
||||
else:
|
||||
curr_interval.start = prev_interval.end
|
||||
i = j
|
||||
|
||||
overlap_start = max(self.event.start_time_ns, curr_interval.start)
|
||||
overlap_end = min(self.event.end_time_ns, curr_interval.end)
|
||||
if overlap_start < overlap_end:
|
||||
overlap_time += overlap_end - overlap_start
|
||||
|
||||
return overlap_time
|
||||
|
||||
|
||||
class BasicEvaluation:
|
||||
def __init__(self, prof: profile) -> None:
|
||||
self.profile = prof
|
||||
self.metrics: dict[EventKey, EventMetrics] = {}
|
||||
self.compute_self_time()
|
||||
self.event_keys = sorted(
|
||||
self.metrics.keys(), key=lambda x: x.event.start_time_ns
|
||||
)
|
||||
self.events = [e.event for e in self.event_keys]
|
||||
self.cuda_events: list[_KinetoEvent] = []
|
||||
self.queue_depth_list = self.compute_queue_depth()
|
||||
self.compute_idle_time()
|
||||
|
||||
def compute_self_time(self) -> None:
|
||||
"""
|
||||
Computes event's self time(total time - time in child ops).
|
||||
"""
|
||||
if self.profile.kineto_results is None:
|
||||
raise AssertionError("kineto_results must not be None")
|
||||
stack = deque(self.profile.kineto_results.experimental_event_tree())
|
||||
|
||||
# standard iterating dfs
|
||||
while stack:
|
||||
curr_event = stack.pop()
|
||||
self_time = curr_event.duration_time_ns
|
||||
for child_event in curr_event.children:
|
||||
self_time -= child_event.duration_time_ns
|
||||
stack.append(child_event)
|
||||
if EventKey(curr_event) in self.metrics:
|
||||
raise AssertionError(
|
||||
f"Duplicate id: {curr_event.id}, {curr_event.name}"
|
||||
)
|
||||
self.metrics[EventKey(curr_event)] = EventMetrics(self_time_ns=self_time)
|
||||
self.metrics[
|
||||
EventKey(curr_event)
|
||||
].duration_time_ns = curr_event.duration_time_ns
|
||||
|
||||
def compute_queue_depth(self):
|
||||
"""
|
||||
Computes queue_depth at each event. This will calculate the queue depth data for
|
||||
All the events in the tree.
|
||||
This will return a list of Interval of queue depth data of cuda launch and kernels.
|
||||
"""
|
||||
if self.profile.kineto_results is None:
|
||||
raise AssertionError("kineto_results must not be None")
|
||||
cuda_event_list = self.profile.kineto_results.events()
|
||||
|
||||
def is_cuda_launch_kernel(e):
|
||||
"""Check if the event is a CUDA launch kernel."""
|
||||
launch_patterns = {
|
||||
"cudaLaunchKernel", # Standard CUDA
|
||||
"cudaLaunchKernelExC", # Extended C
|
||||
"__cudaLaunchKernel", # Internal
|
||||
"cudaLaunchCooperativeKernel", # Collaborative (single-device)
|
||||
"cudaLaunchCooperativeKernelMultiDevice", # Collaborative (multi-devices)
|
||||
}
|
||||
name = str(getattr(e, "name", e))
|
||||
return any(name.startswith(pattern) for pattern in launch_patterns)
|
||||
|
||||
def is_cuda_kernel(e):
|
||||
"""Check if the event is a CUDA runtime kernel."""
|
||||
# Check if the kernel is CUDA
|
||||
if e.device_type() != DeviceType.CUDA:
|
||||
return False
|
||||
|
||||
name = str(getattr(e, "name", e)).lower()
|
||||
|
||||
# Exclude memory operations
|
||||
exclude_patterns = {"mem", "cpy", "alloc", "free"}
|
||||
|
||||
return not any(pattern in name for pattern in exclude_patterns)
|
||||
|
||||
cuda_launch_events = sorted(
|
||||
(e for e in cuda_event_list if is_cuda_launch_kernel(e)),
|
||||
key=lambda x: x.start_ns(),
|
||||
)
|
||||
cuda_kernel_events = sorted(
|
||||
(e for e in cuda_event_list if is_cuda_kernel(e)),
|
||||
key=lambda x: x.start_ns(),
|
||||
)
|
||||
|
||||
self.cuda_events = sorted(
|
||||
cuda_launch_events + cuda_kernel_events, key=lambda x: x.start_ns()
|
||||
)
|
||||
|
||||
kernel_mapping: dict[_KinetoEvent, int] = {}
|
||||
last_mapped_kernel = 0
|
||||
for cuda_launch_event in cuda_launch_events:
|
||||
index = index_of_first_match(
|
||||
cuda_kernel_events,
|
||||
lambda x: x.linked_correlation_id()
|
||||
== cuda_launch_event.linked_correlation_id(),
|
||||
start=last_mapped_kernel,
|
||||
)
|
||||
kernel_mapping[cuda_launch_event] = index
|
||||
last_mapped_kernel = index if index is not None else last_mapped_kernel
|
||||
|
||||
current_kernel_index = 0
|
||||
spawned_kernel_index = -1
|
||||
|
||||
all_events = cuda_launch_events + cuda_kernel_events + self.events
|
||||
|
||||
def new_old_event_comparator(event):
|
||||
if hasattr(event, "start_us"):
|
||||
return event.start_us() * 1000
|
||||
if hasattr(event, "start_ns"):
|
||||
return event.start_ns()
|
||||
if hasattr(event, "start_time_ns"):
|
||||
return event.start_time_ns
|
||||
raise Exception("Unknown Event Type") # noqa: TRY002
|
||||
|
||||
queue_depth_list: list[Interval] = []
|
||||
all_events.sort(key=new_old_event_comparator)
|
||||
for event in all_events:
|
||||
# Find latest cuda kernel event
|
||||
if hasattr(event, "start_us"):
|
||||
start_time = event.start_us() * 1000
|
||||
# pyrefly: ignore [missing-attribute]
|
||||
end_time = (event.start_us() + event.duration_us()) * 1000
|
||||
# Find current spawned cuda kernel event
|
||||
if event in kernel_mapping and kernel_mapping[event] is not None:
|
||||
spawned_kernel_index = kernel_mapping[event]
|
||||
if hasattr(event, "start_ns"):
|
||||
start_time = event.start_ns()
|
||||
end_time = event.start_ns() + event.duration_ns()
|
||||
# Find current spawned cuda kernel event
|
||||
if event in kernel_mapping and kernel_mapping[event] is not None:
|
||||
spawned_kernel_index = kernel_mapping[event]
|
||||
elif hasattr(event, "start_time_ns"):
|
||||
start_time = event.start_time_ns # type: ignore[attr-defined]
|
||||
end_time = event.end_time_ns # type: ignore[attr-defined]
|
||||
|
||||
while (
|
||||
current_kernel_index < len(cuda_kernel_events)
|
||||
and (cuda_kernel_events[current_kernel_index].start_ns()) <= start_time # type: ignore[possibly-undefined]
|
||||
):
|
||||
current_kernel_index += 1
|
||||
current_queue_depth = spawned_kernel_index - current_kernel_index + 1
|
||||
current_queue_depth = max(current_queue_depth, 0)
|
||||
|
||||
if hasattr(event, "start_us") or hasattr(event, "start_ns"):
|
||||
queue_depth_list.append(
|
||||
Interval(start_time, end_time, current_queue_depth) # type: ignore[possibly-undefined]
|
||||
)
|
||||
elif hasattr(event, "start_time_ns"):
|
||||
self.metrics[EventKey(event)].queue_depth = current_queue_depth
|
||||
|
||||
return queue_depth_list
|
||||
|
||||
def compute_idle_time(self) -> None:
|
||||
"""
|
||||
Computes idle time of the profile.
|
||||
"""
|
||||
# Based on queue_depth_list, we can calculate idle time for all the events
|
||||
idle = False
|
||||
idle_start = 0
|
||||
idle_intervals: list[Interval] = []
|
||||
if self.queue_depth_list and self.events:
|
||||
idle_intervals += [
|
||||
Interval(self.events[0].start_time_ns, self.queue_depth_list[0].start),
|
||||
Interval(self.queue_depth_list[-1].end, self.events[-1].end_time_ns),
|
||||
]
|
||||
|
||||
for data_point in self.queue_depth_list:
|
||||
if data_point.queue_depth == 0 and not idle:
|
||||
idle_start = data_point.end
|
||||
idle = True
|
||||
if data_point.queue_depth > 0 and idle:
|
||||
idle_intervals.append(Interval(idle_start, data_point.start))
|
||||
idle = False
|
||||
|
||||
event_list = [e.event for e in self.metrics]
|
||||
for event in event_list:
|
||||
self.metrics[EventKey(event)].idle_time_ns = EventKey(
|
||||
event
|
||||
).intervals_overlap(idle_intervals)
|
||||
|
||||
def rank_events(self, length):
|
||||
"""
|
||||
Filter and Rank the events based on some heuristics:
|
||||
1) Events that are in the falling phase of the queue depth.
|
||||
2) Events that have a high idle_time, self_time difference.
|
||||
|
||||
Parameters:
|
||||
length: The number of events to return.
|
||||
"""
|
||||
|
||||
# Find the interval when qd is falling to 0
|
||||
import torch
|
||||
|
||||
queue_depth_list = list(reversed(self.queue_depth_list))
|
||||
qd_values = [e.queue_depth for e in queue_depth_list]
|
||||
|
||||
bottom_threashold = 0
|
||||
top_threashold = 4
|
||||
decrease_interval = []
|
||||
i = 0
|
||||
while i < len(qd_values):
|
||||
if qd_values[i] > bottom_threashold:
|
||||
i += 1
|
||||
continue
|
||||
for j in range(i + 1, len(qd_values)):
|
||||
# Find next zero and if the max value between them exceeds
|
||||
# the threshold, then we have a falling interval
|
||||
next_minimum_idx = index_of_first_match(
|
||||
qd_values, lambda x: x <= bottom_threashold, start=j
|
||||
)
|
||||
peak_idx = argmax(qd_values, start=j, end=next_minimum_idx)
|
||||
|
||||
# if is a valid peak, we add to list and continue
|
||||
if peak_idx is not None and qd_values[peak_idx] >= top_threashold:
|
||||
decrease_interval.append(
|
||||
Interval(
|
||||
queue_depth_list[peak_idx].start, queue_depth_list[i].start
|
||||
)
|
||||
)
|
||||
i = next_minimum_idx if next_minimum_idx is not None else i
|
||||
break
|
||||
i += 1
|
||||
# Filter out events that are not in the decrease interval
|
||||
event_list = [
|
||||
event
|
||||
for event in self.metrics
|
||||
if event.intervals_overlap(decrease_interval)
|
||||
]
|
||||
if event_list:
|
||||
self_time = torch.tensor(
|
||||
[self.metrics[event].self_time_ns for event in event_list],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
idle_time = torch.tensor(
|
||||
[self.metrics[event].fraction_idle_time for event in event_list],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
normalized_gain = (idle_time - torch.mean(idle_time)) / torch.std(idle_time)
|
||||
normalized_self = (self_time - torch.mean(self_time)) / torch.std(self_time)
|
||||
heuristic_score_list = normalized_gain + 0.6 * normalized_self
|
||||
|
||||
# Sort events by heuristic
|
||||
event_list = [
|
||||
event
|
||||
for _, event in sorted(
|
||||
zip(heuristic_score_list, event_list, strict=True),
|
||||
key=operator.itemgetter(0),
|
||||
reverse=True,
|
||||
)
|
||||
]
|
||||
event_list = event_list[:length]
|
||||
return event_list
|
||||
|
||||
def get_optimizable_events(self, length: int = 1, print_enable: bool = True):
|
||||
event_list = self.rank_events(length)
|
||||
if not print_enable:
|
||||
return event_list
|
||||
output = "Optimizable events:\n" if event_list else "No events to optimize\n"
|
||||
|
||||
output += "\n".join(
|
||||
[
|
||||
f"""{"-" * 80}
|
||||
Event: {event}
|
||||
Source code location: {source_code_location(event.event)}
|
||||
Percentage idle time: {self.metrics[event].fraction_idle_time * 100:.2f}%
|
||||
{"-" * 80}"""
|
||||
for event in event_list
|
||||
]
|
||||
)
|
||||
if print_enable:
|
||||
print(output)
|
||||
return event_list
|
||||
|
||||
|
||||
def index_of_first_match(seq, predicate, start=0, end=None):
|
||||
if end is None or end >= len(seq):
|
||||
end = len(seq)
|
||||
for i in range(start, end):
|
||||
if predicate(seq[i]):
|
||||
return i
|
||||
return None
|
||||
|
||||
|
||||
def argmax(seq, key=lambda x: x, start=0, end=None):
|
||||
seq = seq[start:end]
|
||||
if len(seq) == 0:
|
||||
return None
|
||||
return seq.index(max(seq, key=key)) + start
|
||||
|
||||
|
||||
def source_code_location(event):
|
||||
while event is not None:
|
||||
match = re.search(r"\.py\(.*\)", event.name)
|
||||
if match is None:
|
||||
event = event.parent
|
||||
continue
|
||||
return event.name
|
||||
return "No source code location found"
|
||||
|
||||
|
||||
# Provide an OSS workaround for cudagraphs + CUPTI issue
|
||||
# https://github.com/pytorch/pytorch/issues/75504
|
||||
# TODO(dberard) - deprecate / remove workaround for CUDA >= 12, when
|
||||
# we stop supporting older CUDA versions.
|
||||
def _init_for_cuda_graphs() -> None:
|
||||
from torch.autograd.profiler import profile
|
||||
|
||||
with profile():
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class TimelineEvent:
|
||||
"""Represents an event in the profiler timeline."""
|
||||
|
||||
timestamp: int
|
||||
event_type: Literal["start", "end", "regular"]
|
||||
marker_type: Literal["filename", "node"] | None
|
||||
identifier: str | int | None
|
||||
event: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContextStackEntry:
|
||||
"""Represents a context (filename or node) in the stack."""
|
||||
|
||||
context_type: Literal["filename", "node"]
|
||||
identifier: str | int
|
||||
metadata: dict | None
|
||||
tid: int | None = None # Thread ID associated with this context
|
||||
|
||||
|
||||
def map_recorded_events_to_aten_ops_with_stack_trace(traced_data):
|
||||
"""
|
||||
Maps recorded profiler events to their corresponding fx nodes and adds stack traces.
|
||||
|
||||
Builds a timeline of all events (regular ops and FX markers for filenames/nodes),
|
||||
sorts by timestamp, then processes chronologically while maintaining a context stack of active
|
||||
filename/node scopes. Regular events are augmented with stack traces and node names from the
|
||||
innermost active context. Runtime is O(n log n) for n events.
|
||||
|
||||
Args:
|
||||
traced_data: Json of profiler events from Chrome trace
|
||||
|
||||
Returns:
|
||||
Dict mapping recorded event names to their aten operations with added stack traces
|
||||
"""
|
||||
from torch.fx.traceback import _FX_METADATA_REGISTRY
|
||||
|
||||
trace_events = traced_data.get("traceEvents", [])
|
||||
|
||||
# Create event timeline
|
||||
event_timeline: list[TimelineEvent] = []
|
||||
|
||||
def is_fx_marker_event(event):
|
||||
return (
|
||||
event.get("cat") == "cpu_op"
|
||||
and event.get("name", "").startswith("## ")
|
||||
and event.get("name", "").endswith(" ##")
|
||||
)
|
||||
|
||||
def append_fx_marker_event(event_type, identifier, event):
|
||||
start_ts = event["ts"]
|
||||
end_ts = start_ts + event["dur"]
|
||||
event_timeline.append(
|
||||
TimelineEvent(start_ts, "start", event_type, identifier, event)
|
||||
)
|
||||
event_timeline.append(
|
||||
TimelineEvent(end_ts, "end", event_type, identifier, event)
|
||||
)
|
||||
|
||||
for event in trace_events:
|
||||
if "ts" not in event or "dur" not in event:
|
||||
continue
|
||||
|
||||
if is_fx_marker_event(event):
|
||||
content = event["name"][3:-3]
|
||||
|
||||
if content.endswith(".py"):
|
||||
append_fx_marker_event("filename", content, event)
|
||||
else:
|
||||
try:
|
||||
node_index = int(content)
|
||||
except ValueError:
|
||||
pass
|
||||
append_fx_marker_event("node", node_index, event) # type: ignore[possibly-undefined]
|
||||
|
||||
else:
|
||||
# Regular event that needs augmentation
|
||||
start_ts = event["ts"]
|
||||
event_timeline.append(TimelineEvent(start_ts, "regular", None, None, event))
|
||||
|
||||
# Sort by timestamp
|
||||
event_timeline.sort(key=lambda x: x.timestamp)
|
||||
|
||||
# Process events in chronological order with a stack
|
||||
context_stack: list[ContextStackEntry] = []
|
||||
|
||||
# Invariant: all start event has a corresponding end event
|
||||
for timeline_event in event_timeline:
|
||||
match timeline_event.event_type:
|
||||
case "start":
|
||||
if timeline_event.identifier is None:
|
||||
raise AssertionError("identifier must not be None for start event")
|
||||
|
||||
if timeline_event.marker_type == "filename":
|
||||
if not isinstance(timeline_event.identifier, str):
|
||||
raise AssertionError(
|
||||
f"identifier must be str for filename marker, "
|
||||
f"got {type(timeline_event.identifier).__name__}"
|
||||
)
|
||||
# Push filename context - query metadata registry on-demand
|
||||
metadata = _FX_METADATA_REGISTRY.get(timeline_event.identifier)
|
||||
tid = timeline_event.event.get("tid")
|
||||
context_stack.append(
|
||||
ContextStackEntry(
|
||||
"filename", timeline_event.identifier, metadata, tid
|
||||
)
|
||||
)
|
||||
elif timeline_event.marker_type == "node":
|
||||
# Find the current filename from stack
|
||||
current_file_metadata = None
|
||||
tid = timeline_event.event.get("tid")
|
||||
for ctx_entry in reversed(context_stack):
|
||||
if (
|
||||
ctx_entry.context_type == "filename"
|
||||
and ctx_entry.tid == tid
|
||||
):
|
||||
current_file_metadata = ctx_entry.metadata
|
||||
break
|
||||
|
||||
if current_file_metadata:
|
||||
node_metadata = current_file_metadata.get("node_metadata", {})
|
||||
if timeline_event.identifier in node_metadata:
|
||||
node_meta: dict | None = node_metadata[
|
||||
timeline_event.identifier
|
||||
]
|
||||
context_stack.append(
|
||||
ContextStackEntry(
|
||||
"node", timeline_event.identifier, node_meta, tid
|
||||
)
|
||||
)
|
||||
|
||||
case "end":
|
||||
# Pop from stack - search backwards to find matching context
|
||||
for i in range(len(context_stack) - 1, -1, -1):
|
||||
ctx_entry = context_stack[i]
|
||||
if (
|
||||
timeline_event.marker_type == ctx_entry.context_type
|
||||
and timeline_event.identifier == ctx_entry.identifier
|
||||
):
|
||||
context_stack.pop(i)
|
||||
break
|
||||
|
||||
case "regular":
|
||||
# Apply metadata from current context stack
|
||||
# Find the most specific context (node takes precedence over filename)
|
||||
# Only augment events with the same tid as the file/node event matched
|
||||
current_stack_trace = None
|
||||
current_node_name = None
|
||||
event_tid = timeline_event.event.get("tid")
|
||||
|
||||
for ctx_entry in reversed(context_stack):
|
||||
# Only apply metadata from contexts with matching tid
|
||||
if ctx_entry.tid == event_tid:
|
||||
if ctx_entry.context_type == "node" and ctx_entry.metadata:
|
||||
current_stack_trace = ctx_entry.metadata.get(
|
||||
"stack_trace", "No model stack trace available"
|
||||
)
|
||||
current_node_name = ctx_entry.metadata.get("name", "")
|
||||
# Do we want to only attach the stack trace of the lowest node or stack trace of all nodes
|
||||
# if nodes are nested, e.g. in nested graph modules
|
||||
break
|
||||
|
||||
# Augment the event
|
||||
if current_stack_trace or current_node_name:
|
||||
args = timeline_event.event.setdefault("args", {})
|
||||
if current_stack_trace:
|
||||
args["stack_trace"] = current_stack_trace
|
||||
if current_node_name:
|
||||
args["node_name"] = current_node_name
|
||||
@@ -0,0 +1,81 @@
|
||||
# mypy: allow-untyped-defs
|
||||
from contextlib import contextmanager
|
||||
from typing import NoReturn
|
||||
|
||||
|
||||
try:
|
||||
from torch._C import _itt
|
||||
except ImportError:
|
||||
|
||||
class _ITTStub:
|
||||
@staticmethod
|
||||
def _fail(*args, **kwargs) -> NoReturn:
|
||||
raise RuntimeError(
|
||||
"ITT functions not installed. Are you sure you have a ITT build?"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def is_available() -> bool:
|
||||
return False
|
||||
|
||||
rangePush = _fail
|
||||
rangePop = _fail
|
||||
mark = _fail
|
||||
|
||||
_itt = _ITTStub() # type: ignore[assignment]
|
||||
|
||||
|
||||
__all__ = ["is_available", "range_push", "range_pop", "mark", "range"]
|
||||
|
||||
|
||||
def is_available():
|
||||
"""
|
||||
Check if ITT feature is available or not
|
||||
"""
|
||||
return _itt.is_available()
|
||||
|
||||
|
||||
def range_push(msg):
|
||||
"""
|
||||
Pushes a range onto a stack of nested range span. Returns zero-based
|
||||
depth of the range that is started.
|
||||
|
||||
Arguments:
|
||||
msg (str): ASCII message to associate with range
|
||||
"""
|
||||
return _itt.rangePush(msg)
|
||||
|
||||
|
||||
def range_pop():
|
||||
"""
|
||||
Pops a range off of a stack of nested range spans. Returns the
|
||||
zero-based depth of the range that is ended.
|
||||
"""
|
||||
return _itt.rangePop()
|
||||
|
||||
|
||||
def mark(msg):
|
||||
"""
|
||||
Describe an instantaneous event that occurred at some point.
|
||||
|
||||
Arguments:
|
||||
msg (str): ASCII message to associate with the event.
|
||||
"""
|
||||
return _itt.mark(msg)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def range(msg, *args, **kwargs):
|
||||
"""
|
||||
Context manager / decorator that pushes an ITT range at the beginning
|
||||
of its scope, and pops it at the end. If extra arguments are given,
|
||||
they are passed as arguments to msg.format().
|
||||
|
||||
Args:
|
||||
msg (str): message to associate with the range
|
||||
"""
|
||||
range_push(msg.format(*args, **kwargs))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
range_pop()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,20 @@
|
||||
import os
|
||||
import site
|
||||
import sys
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _prefix_regex() -> list[str]:
|
||||
raw_paths = (
|
||||
site.getsitepackages()
|
||||
+ sys.path
|
||||
+ [site.getuserbase()]
|
||||
+ [site.getusersitepackages()]
|
||||
+ [os.path.dirname(os.path.dirname(torch.__file__))]
|
||||
)
|
||||
|
||||
path_prefixes = sorted({os.path.abspath(i) for i in raw_paths}, reverse=True)
|
||||
if not all(isinstance(i, str) for i in path_prefixes):
|
||||
raise AssertionError("all path_prefixes must be strings")
|
||||
return [i + os.sep for i in path_prefixes]
|
||||
Reference in New Issue
Block a user