From ca1e502a655993339715c9f4e48349a634fe2906 Mon Sep 17 00:00:00 2001 From: Jassani Date: Wed, 6 May 2026 10:24:43 -0400 Subject: [PATCH 01/10] PerfModel: add op-name -> category registry as v2 of categorize_torch_op Adds a flat ``op_name -> category`` registry alongside the existing ``categorize_torch_op`` if/elif chain. The registry is auto-populated from ``op_to_perf_model_class_map`` + ``dict_base_class2category`` and topped up with explicit overrides for category-only ops (e.g. SSM_bwd, MoE_aux, CONV variants without a perf model) that the legacy chain handled inline. A small regex layer (``triton``, ``record_param_comms``) and a kernel-name fallback (``void at::native`` probes) preserve the rest of the legacy chain's behavior. This PR is purely additive: - The legacy ``categorize_torch_op`` is unchanged. - The new ``categorize_torch_op_v2`` runs in parallel. - ``test_categorize_torch_op_parity`` asserts v1 and v2 return the same category for every name reachable through perf-model registration, every name hardcoded inside the legacy chain, the pattern-matched names, "other" cases, and the kernel-name fallback. Motivation: today categorization is coupled to perf-model availability (category derived from the perf-model class hierarchy), which forces a new perf-model class for every category-only op and pushes hardcoded name lists into the categorizer function. The registry breaks that coupling so subsequent PRs can: - Delete the if/elif chain (PR B). - Add a ``get_op_categories`` extension hook so category-only entries don't require authoring a perf model (PR B). - Add a drift report for ops landing in ``"other"`` (PR C). No user-visible behavior change in this PR. Co-authored-by: Cursor --- TraceLens/PerfModel/op_categories.py | 250 +++++++++++++++++++++++ TraceLens/PerfModel/torch_op_mapping.py | 27 +++ tests/test_categorize_torch_op_parity.py | 197 ++++++++++++++++++ 3 files changed, 474 insertions(+) create mode 100644 TraceLens/PerfModel/op_categories.py create mode 100644 tests/test_categorize_torch_op_parity.py diff --git a/TraceLens/PerfModel/op_categories.py b/TraceLens/PerfModel/op_categories.py new file mode 100644 index 000000000..dd562510f --- /dev/null +++ b/TraceLens/PerfModel/op_categories.py @@ -0,0 +1,250 @@ +############################################################################### +# Copyright (c) 2024 - 2025 Advanced Micro Devices, Inc. All rights reserved. +# +# See LICENSE for license information. +############################################################################### + +""" +Registry-based categorization of CPU torch ops. + +This module provides a flat ``op_name -> category`` registry plus a small set +of fallback patterns. It is the v2 replacement for the if/elif chain in +``categorize_torch_op``. PR A introduces ``categorize_torch_op_v2`` alongside +the legacy implementation; ``test_categorize_torch_op_parity`` asserts both +produce the same output for every reachable name. + +Subsequent PRs (referenced from the meta-issue): + + - PR B: delete the if/elif chain, move the remaining hardcoded names into + ``LEGACY_CATEGORIZE_EXTRAS``, and add a ``get_op_categories`` extension + hook so category-only entries don't require a perf model. + - PR C: add ``report_uncategorized_ops`` for drift detection. +""" + +from __future__ import annotations + +import re +from typing import Iterable, Mapping, Optional + +# --------------------------------------------------------------------------- +# Legacy SDPA backward names +# +# Replicated from the hardcoded list inside ``categorize_torch_op`` so v2 +# produces identical output for backward-attention ops that were previously +# detected by name rather than by the ``_backward`` suffix. +# +# Note: in the legacy chain this list was only consulted for ops already in +# ``dict_cat2names["SDPA"]``. Names below that aren't perf-modeled (e.g. +# ``FlashAttnFuncBackward``) currently resolve to ``"other"`` in v1 and v2 +# preserves that. PR B can decide whether to lift that restriction. +# --------------------------------------------------------------------------- +_LEGACY_SDPA_BWD_NAMES = frozenset( + { + "FlashAttnFuncBackward", + "FusedAttnFuncBackward", + "flash_attn::_flash_attn_backward", + "flash_attn::_flash_attn_varlen_backward", + "aten::_scaled_dot_product_cudnn_attention_backward", + "aten::_scaled_dot_product_efficient_attention_backward", + "aten::_scaled_dot_product_flash_attention_backward", + "aiter::_flash_attn_backward", + "aiter::wrapper_fmha_v3_bwd", + "aiter::mha_bwd", + } +) + + +# --------------------------------------------------------------------------- +# Explicit ``op_name -> category`` overrides +# +# Each entry corresponds to a hardcoded list inside the legacy +# ``categorize_torch_op`` chain. Once PR B deletes that chain these become +# the only path for category-only ops (no perf model required). +# --------------------------------------------------------------------------- +LEGACY_CATEGORIZE_EXTRAS: dict[str, str] = { + # CONV ops not present in op_to_perf_model_class_map + "aten::miopen_convolution": "CONV_fwd", + "aten::cudnn_convolution": "CONV_fwd", + # SSM + "MambaSplitConv1dScanCombinedFnBackward": "SSM_bwd", + "DaoAILab::_causal_conv1d_bwd_cpp": "SSM_bwd", + # MoE_comm forward extras (not perf-modeled) + "TokenPermuteMaskMap": "MoE_comm_fwd", + "_OperationFuserAutogradFunction": "MoE_comm_fwd", + # MoE_comm backward + "MoEDispatchBackward": "MoE_comm_bwd", + "MoECombineBackward": "MoE_comm_bwd", + "TokenPermuteMaskMapBackward": "MoE_comm_bwd", + "_OperationFuserAutogradFunctionBackward": "MoE_comm_bwd", + # RoPE / CrossEntropy backward + "FusedRoPEFuncBackward": "RoPE_bwd", + "CrossEntropyFunctionBackward": "CrossEntropy_bwd", + # MoE auxiliary ops (sorting / topk) + "aiter::moe_sorting_fwd": "MoE_aux", + "aiter::moe_sorting_opus_fwd": "MoE_aux", + "aiter::moe_align_block_size": "MoE_aux", + "_moe_C::moe_align_block_size": "MoE_aux", + "aiter::fused_moe_->_fused_dynamic_mxfp4_quant_moe_sort_kernel (Synthetic Op)": "MoE_aux", + "aiter::moe_sum": "MoE_aux", + "aiter::topk_softmax": "MoE_aux", + "aiter::topk_softmax_asm": "MoE_aux", + "aiter::topk_sigmoid": "MoE_aux", + "aiter::biased_grouped_topk_hip": "MoE_aux", + "aiter::grouped_topk": "MoE_aux", + "aiter::moe_fused_gate": "MoE_aux", + # InferenceAttention extras (KV-cache writes) + "_C_cache_ops::reshape_and_cache_flash": "InferenceAttention", + "_C_cache_ops::concat_and_cache_mla": "InferenceAttention", +} + + +# --------------------------------------------------------------------------- +# Patterns evaluated only when no exact registry match is found, in declared +# order. ``re.match`` is implicitly anchored at the start of the string. +# --------------------------------------------------------------------------- +OP_CATEGORY_PATTERNS: list[tuple[re.Pattern, str]] = [ + (re.compile(r"triton"), "triton"), + (re.compile(r"record_param_comms"), "record_param_comms"), +] + + +# --------------------------------------------------------------------------- +# Kernel-name fallback rules +# +# Applied only to GPU kernel names starting with ``void at::native``. The +# first substring match wins. Mirrors the kernel-name probe at the bottom +# of the legacy ``categorize_torch_op``. +# --------------------------------------------------------------------------- +_KERNEL_NAME_PREFIX = "void at::native" +_KERNEL_NAME_FALLBACK_RULES: tuple[tuple[str, str], ...] = ( + ("elementwise", "elementwise"), + ("reduce", "reduce"), + ("multi_tensor_apply", "multi_tensor_apply"), +) + + +def _kernel_name_fallback(row) -> Optional[str]: + kernel_details = row.get("kernel_details") + if not kernel_details: + return None + kernel_name = kernel_details[0].get("name", "") + if not kernel_name.startswith(_KERNEL_NAME_PREFIX): + return None + for needle, category in _KERNEL_NAME_FALLBACK_RULES: + if needle in kernel_name: + return category + return None + + +def _resolve_base_category(op_name: str, base_category: str) -> str: + """Map ``(op_name, base-class category)`` to the final category emitted by v1. + + The legacy chain applied a handful of fwd/bwd splits and category renames + that aren't expressible by walking the perf-model class hierarchy alone. + This helper centralises those rules. + """ + if base_category == "SDPA": + if op_name.endswith("_backward") or op_name in _LEGACY_SDPA_BWD_NAMES: + return "SDPA_bwd" + return "SDPA_fwd" + if base_category == "Normalization": + if op_name.endswith("_backward") or op_name.endswith("Backward"): + return "NORM_bwd" + return "NORM_fwd" + if base_category == "CONV": + if op_name.endswith("_backward") or op_name.endswith("Backward"): + return "CONV_bwd" + return "CONV_fwd" + if base_category == "SSM": + return "SSM_fwd" + if base_category == "MoE_comm": + return "MoE_comm_fwd" + if base_category == "RoPE": + return "RoPE_fwd" + if base_category == "CrossEntropy": + return "CrossEntropy_fwd" + if base_category in ("BinaryElementwise", "UnaryElementwise"): + return "elementwise" + if base_category == "Reduce": + return "reduce" + return base_category + + +def build_op_category_registry( + op_to_perf_model_class_map: Mapping[str, type], + dict_base_class2category: Mapping[type, str], + extras: Optional[Mapping[str, str]] = None, +) -> dict[str, str]: + """Construct the flat ``op_name -> category`` registry. + + The registry is built in two passes: + + 1. Walk every entry in ``op_to_perf_model_class_map``. Each perf-model + class contributes exactly one ``(op_name, category)`` pair, derived + from its single base class via ``dict_base_class2category`` and then + passed through :func:`_resolve_base_category` to apply fwd/bwd splits. + + 2. Apply ``extras`` as overrides. Use this for category-only ops + (no perf model) and for ops where the legacy chain disagrees with + the auto-derived category. + + Parameters + ---------- + op_to_perf_model_class_map + Mapping from op name to perf-model class. + dict_base_class2category + Mapping from perf-model base class to category label. + extras + Optional ``op_name -> category`` overrides. + + Returns + ------- + dict[str, str] + ``op_name -> final_category`` lookup table. + """ + registry: dict[str, str] = {} + + for op_name, perf_model_class in op_to_perf_model_class_map.items(): + base_classes = perf_model_class.__bases__ + if len(base_classes) != 1: + continue + base_category = dict_base_class2category.get(base_classes[0]) + if base_category is None: + continue + registry[op_name] = _resolve_base_category(op_name, base_category) + + if extras: + registry.update(extras) + + return registry + + +def categorize_torch_op_v2( + row, + registry: Mapping[str, str], + patterns: Iterable[tuple[re.Pattern, str]] = OP_CATEGORY_PATTERNS, +) -> str: + """Return the category for ``row`` using the flat registry plus patterns. + + Resolution order: + + 1. Exact name match in ``registry``. + 2. First pattern in ``patterns`` whose regex matches ``row["name"]``. + 3. Kernel-name fallback for ``void at::native`` GPU kernels. + 4. ``"other"``. + """ + name = row["name"] + + cat = registry.get(name) + if cat is not None: + return cat + + for pattern, category in patterns: + if pattern.match(name): + return category + + fallback = _kernel_name_fallback(row) + if fallback is not None: + return fallback + + return "other" diff --git a/TraceLens/PerfModel/torch_op_mapping.py b/TraceLens/PerfModel/torch_op_mapping.py index 8bded969a..e4aa33b84 100644 --- a/TraceLens/PerfModel/torch_op_mapping.py +++ b/TraceLens/PerfModel/torch_op_mapping.py @@ -7,6 +7,11 @@ from . import perf_model from collections import defaultdict from .extensions import get_pseudo_op_mappings, get_pseudo_op_categories +from .op_categories import ( + LEGACY_CATEGORIZE_EXTRAS, + build_op_category_registry, + categorize_torch_op_v2 as _categorize_v2, +) op_to_perf_model_class_map = { "aten::mm": perf_model.aten_mm, @@ -399,3 +404,25 @@ def categorize_torch_op(row): return "multi_tensor_apply" # if none of the above cases match, return 'other' return "other" + + +# --------------------------------------------------------------------------- +# Registry-based v2 categorizer (introduced in PR A; runs alongside the +# legacy ``categorize_torch_op`` chain above). The parity test +# ``tests/test_categorize_torch_op_parity.py`` asserts both produce the same +# output for every reachable name. PR B will swap v2 in and delete v1. +# --------------------------------------------------------------------------- +OP_CATEGORY_REGISTRY = build_op_category_registry( + op_to_perf_model_class_map, + dict_base_class2category, + extras=LEGACY_CATEGORIZE_EXTRAS, +) + + +def categorize_torch_op_v2(row): + """Registry-based replacement for :func:`categorize_torch_op`. + + Equivalent in behavior to the legacy if/elif chain (verified by the + parity test). See ``op_categories.py`` for the design rationale. + """ + return _categorize_v2(row, OP_CATEGORY_REGISTRY) diff --git a/tests/test_categorize_torch_op_parity.py b/tests/test_categorize_torch_op_parity.py new file mode 100644 index 000000000..6122a0004 --- /dev/null +++ b/tests/test_categorize_torch_op_parity.py @@ -0,0 +1,197 @@ +############################################################################### +# Copyright (c) 2024 - 2025 Advanced Micro Devices, Inc. All rights reserved. +# +# See LICENSE for license information. +############################################################################### + +""" +Parity test for :func:`categorize_torch_op_v2` (registry-based) against +:func:`categorize_torch_op` (legacy if/elif chain). + +PR A introduces v2 alongside v1 without changing v1's behavior. The single +acceptance criterion is that v2 produces the same output as v1 for every +reachable input. PR B will delete v1; that PR's safety relies on this test. +""" + +import pytest + +from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, + categorize_torch_op, + categorize_torch_op_v2, + op_to_perf_model_class_map, +) + + +def _row(name, kernel_details=None): + return {"name": name, "kernel_details": kernel_details or []} + + +# --------------------------------------------------------------------------- +# Coverage set 1: every name reachable through perf-model registration +# (op_to_perf_model_class_map, including extension-contributed entries). +# --------------------------------------------------------------------------- +PERF_MODEL_NAMES = sorted(op_to_perf_model_class_map.keys()) + + +# --------------------------------------------------------------------------- +# Coverage set 2: names hardcoded inside the legacy categorize_torch_op +# if/elif chain. These are the names PR A had to thread into +# LEGACY_CATEGORIZE_EXTRAS to preserve parity. +# --------------------------------------------------------------------------- +LEGACY_HARDCODED_NAMES = [ + # CONV_fwd + "aten::convolution", + "aten::miopen_convolution", + "aten::cudnn_convolution", + "ConvBias_", + "ConvBiasReLU_", + # CONV_bwd + "aten::convolution_backward", + "ConvBias_Backward", + "ConvBiasReLU_Backward", + # SDPA backward names (some perf-modeled, some unreachable in v1) + "FlashAttnFuncBackward", + "FusedAttnFuncBackward", + "flash_attn::_flash_attn_backward", + "flash_attn::_flash_attn_varlen_backward", + "aten::_scaled_dot_product_cudnn_attention_backward", + "aten::_scaled_dot_product_efficient_attention_backward", + "aten::_scaled_dot_product_flash_attention_backward", + "aiter::_flash_attn_backward", + "aiter::wrapper_fmha_v3_bwd", + "aiter::mha_bwd", + # SSM + "MambaSplitConv1dScanCombinedFn", + "MambaSplitConv1dScanCombinedFnBackward", + "DaoAILab::_causal_conv1d_bwd_cpp", + # MoE_comm + "TokenPermuteMaskMap", + "_OperationFuserAutogradFunction", + "MoEDispatchBackward", + "MoECombineBackward", + "TokenPermuteMaskMapBackward", + "_OperationFuserAutogradFunctionBackward", + # RoPE / CrossEntropy backward + "FusedRoPEFuncBackward", + "CrossEntropyFunctionBackward", + # MoE auxiliary + "aiter::moe_sorting_fwd", + "aiter::moe_sorting_opus_fwd", + "aiter::moe_align_block_size", + "_moe_C::moe_align_block_size", + "aiter::fused_moe_->_fused_dynamic_mxfp4_quant_moe_sort_kernel (Synthetic Op)", + "aiter::moe_sum", + "aiter::topk_softmax", + "aiter::topk_softmax_asm", + "aiter::topk_sigmoid", + "aiter::biased_grouped_topk_hip", + "aiter::grouped_topk", + "aiter::moe_fused_gate", + # InferenceAttention extras + "_C_cache_ops::reshape_and_cache_flash", + "_C_cache_ops::concat_and_cache_mla", +] + + +# --------------------------------------------------------------------------- +# Coverage set 3: names matched by patterns (triton, record_param_comms). +# --------------------------------------------------------------------------- +PATTERN_NAMES = [ + "triton", + "triton_softmax_kernel", + "triton_per_token_quant_fp8", + "record_param_comms", + "record_param_comms_alltoall_base", +] + + +# --------------------------------------------------------------------------- +# Coverage set 4: unknown names that should resolve to "other" (or to +# a kernel-name fallback). +# --------------------------------------------------------------------------- +OTHER_NAMES = [ + "completely::unknown::op", + "no_match_at_all", + "Some::Random::Op", +] + + +# --------------------------------------------------------------------------- +# Coverage set 5: kernel-name fallback inputs. +# --------------------------------------------------------------------------- +KERNEL_FALLBACK_CASES = [ + # (kernel_name, expected_category) + ("void at::native::elementwise_kernel<8, 4>(...)", "elementwise"), + ("void at::native::reduce_kernel<512, 4>(...)", "reduce"), + ("void at::native::multi_tensor_apply_kernel<...>(...)", "multi_tensor_apply"), + # Native kernel that matches no needle -> "other" + ("void at::native::unrelated_kernel(...)", "other"), + # Non-native kernel -> "other" + ("some_user_kernel(...)", "other"), +] + + +# --------------------------------------------------------------------------- +# Parametrized parity assertions +# --------------------------------------------------------------------------- +@pytest.mark.parametrize("name", PERF_MODEL_NAMES) +def test_parity_perf_model_names(name): + row = _row(name) + assert categorize_torch_op_v2(row) == categorize_torch_op(row) + + +@pytest.mark.parametrize("name", LEGACY_HARDCODED_NAMES) +def test_parity_legacy_hardcoded_names(name): + row = _row(name) + assert categorize_torch_op_v2(row) == categorize_torch_op(row) + + +@pytest.mark.parametrize("name", PATTERN_NAMES) +def test_parity_pattern_names(name): + row = _row(name) + assert categorize_torch_op_v2(row) == categorize_torch_op(row) + + +@pytest.mark.parametrize("name", OTHER_NAMES) +def test_parity_other_names(name): + row = _row(name) + assert categorize_torch_op_v2(row) == categorize_torch_op(row) + + +@pytest.mark.parametrize("kernel_name,expected", KERNEL_FALLBACK_CASES) +def test_parity_kernel_fallback(kernel_name, expected): + row = _row("unknown_op_name", kernel_details=[{"name": kernel_name}]) + v1 = categorize_torch_op(row) + v2 = categorize_torch_op_v2(row) + assert v1 == v2 == expected, f"v1={v1!r}, v2={v2!r}, expected={expected!r}" + + +# --------------------------------------------------------------------------- +# Sanity checks on the registry shape +# --------------------------------------------------------------------------- +def test_registry_built_at_import(): + assert isinstance(OP_CATEGORY_REGISTRY, dict) + assert len(OP_CATEGORY_REGISTRY) > 0 + + +def test_registry_contains_expected_entries(): + # Auto-derived from a perf-model base class + assert OP_CATEGORY_REGISTRY["aten::mm"] == "GEMM" + # Auto-derived NORM split + assert OP_CATEGORY_REGISTRY["aten::layer_norm"] == "NORM_fwd" + assert OP_CATEGORY_REGISTRY["aten::layer_norm_backward"] == "NORM_bwd" + # Auto-derived SDPA fwd vs explicit-list bwd + assert OP_CATEGORY_REGISTRY["FlashAttnFunc"] == "SDPA_fwd" + assert OP_CATEGORY_REGISTRY["aiter::wrapper_fmha_v3_bwd"] == "SDPA_bwd" + # Extras-only entries + assert OP_CATEGORY_REGISTRY["aten::miopen_convolution"] == "CONV_fwd" + assert OP_CATEGORY_REGISTRY["DaoAILab::_causal_conv1d_bwd_cpp"] == "SSM_bwd" + assert OP_CATEGORY_REGISTRY["aiter::topk_softmax"] == "MoE_aux" + + +def test_registry_covers_every_perf_model_op(): + # The builder must produce a category for every perf-modeled op (i.e. no + # silent drops). PR B relies on this invariant. + missing = [n for n in op_to_perf_model_class_map if n not in OP_CATEGORY_REGISTRY] + assert missing == [], f"perf-modeled ops missing from registry: {missing}" From 54d49d36309f0a2eaa9285fce3723b124b6b3085 Mon Sep 17 00:00:00 2001 From: Jassani Date: Wed, 6 May 2026 12:15:50 -0400 Subject: [PATCH 02/10] PerfModel: replace torch op categorizer with registry Replaces the legacy if/elif chain in categorize_torch_op with the registry-backed categorizer directly. Keeps dict_cat2names as the compatibility view for legacy report sheets, while OP_CATEGORY_REGISTRY becomes the source of truth for final op labels. This also splits GroupedGEMM out of GEMM, fixes SDPA backward names that previously fell through unless dict_cat2names was mutated, and wires perf_model_extension, op_category_extension, and dict_cat2names_extension into the registry. Focused tests cover expected categories, kernel fallback behavior, extension registration, grouped GEMM categorization, and the existing pseudo-op extension compatibility cases. Co-authored-by: Cursor --- TraceLens/PerfModel/op_categories.py | 325 +++++++++++------- TraceLens/PerfModel/torch_op_mapping.py | 211 ++---------- .../Reporting/generate_perf_report_pytorch.py | 40 ++- .../generate_perf_report_pytorch_inference.py | 40 ++- tests/test_categorize_torch_op_parity.py | 197 ----------- tests/test_primus_op_categorization.py | 8 +- tests/test_primus_turbo_grouped_gemm.py | 4 +- .../test_torch_op_categorization_registry.py | 138 ++++++++ 8 files changed, 429 insertions(+), 534 deletions(-) delete mode 100644 tests/test_categorize_torch_op_parity.py create mode 100644 tests/test_torch_op_categorization_registry.py diff --git a/TraceLens/PerfModel/op_categories.py b/TraceLens/PerfModel/op_categories.py index dd562510f..f64372e0f 100644 --- a/TraceLens/PerfModel/op_categories.py +++ b/TraceLens/PerfModel/op_categories.py @@ -7,38 +7,19 @@ """ Registry-based categorization of CPU torch ops. -This module provides a flat ``op_name -> category`` registry plus a small set -of fallback patterns. It is the v2 replacement for the if/elif chain in -``categorize_torch_op``. PR A introduces ``categorize_torch_op_v2`` alongside -the legacy implementation; ``test_categorize_torch_op_parity`` asserts both -produce the same output for every reachable name. - -Subsequent PRs (referenced from the meta-issue): - - - PR B: delete the if/elif chain, move the remaining hardcoded names into - ``LEGACY_CATEGORIZE_EXTRAS``, and add a ``get_op_categories`` extension - hook so category-only entries don't require a perf model. - - PR C: add ``report_uncategorized_ops`` for drift detection. +The source of truth is a flat ``op_name -> final category`` registry used by +``categorize_torch_op``. ``dict_cat2names`` remains as a compatibility view for +legacy per-category sheets such as ``GEMM``, ``CONV_fwd`` / ``CONV_bwd``, and +``GroupedGEMM_fwd`` / ``GroupedGEMM_bwd``. """ -from __future__ import annotations - import re -from typing import Iterable, Mapping, Optional +from collections import defaultdict +from typing import DefaultDict, Dict, Iterable, List, Mapping, MutableMapping +from typing import Optional, Pattern, Tuple -# --------------------------------------------------------------------------- -# Legacy SDPA backward names -# -# Replicated from the hardcoded list inside ``categorize_torch_op`` so v2 -# produces identical output for backward-attention ops that were previously -# detected by name rather than by the ``_backward`` suffix. -# -# Note: in the legacy chain this list was only consulted for ops already in -# ``dict_cat2names["SDPA"]``. Names below that aren't perf-modeled (e.g. -# ``FlashAttnFuncBackward``) currently resolve to ``"other"`` in v1 and v2 -# preserves that. PR B can decide whether to lift that restriction. -# --------------------------------------------------------------------------- -_LEGACY_SDPA_BWD_NAMES = frozenset( + +SDPA_BWD_OPS = frozenset( { "FlashAttnFuncBackward", "FusedAttnFuncBackward", @@ -54,32 +35,32 @@ ) -# --------------------------------------------------------------------------- -# Explicit ``op_name -> category`` overrides -# -# Each entry corresponds to a hardcoded list inside the legacy -# ``categorize_torch_op`` chain. Once PR B deletes that chain these become -# the only path for category-only ops (no perf model required). -# --------------------------------------------------------------------------- -LEGACY_CATEGORIZE_EXTRAS: dict[str, str] = { - # CONV ops not present in op_to_perf_model_class_map +OP_CATEGORY_OVERRIDES: Dict[str, str] = { + # CONV ops not present in op_to_perf_model_class_map. "aten::miopen_convolution": "CONV_fwd", "aten::cudnn_convolution": "CONV_fwd", - # SSM + # SDPA backward ops that do not all have perf models in core TraceLens. + "FlashAttnFuncBackward": "SDPA_bwd", + "FusedAttnFuncBackward": "SDPA_bwd", + "aten::_scaled_dot_product_cudnn_attention_backward": "SDPA_bwd", + "aten::_scaled_dot_product_efficient_attention_backward": "SDPA_bwd", + "aten::_scaled_dot_product_flash_attention_backward": "SDPA_bwd", + # SSM / Mamba category-only backward ops. "MambaSplitConv1dScanCombinedFnBackward": "SSM_bwd", "DaoAILab::_causal_conv1d_bwd_cpp": "SSM_bwd", - # MoE_comm forward extras (not perf-modeled) + # MoE communication category-only ops. "TokenPermuteMaskMap": "MoE_comm_fwd", + # Observed in MoE token-routing traces; tracked separately because the name + # itself is generic and may not always imply MoE communication. "_OperationFuserAutogradFunction": "MoE_comm_fwd", - # MoE_comm backward "MoEDispatchBackward": "MoE_comm_bwd", "MoECombineBackward": "MoE_comm_bwd", "TokenPermuteMaskMapBackward": "MoE_comm_bwd", "_OperationFuserAutogradFunctionBackward": "MoE_comm_bwd", - # RoPE / CrossEntropy backward + # RoPE / CrossEntropy category-only backward ops. "FusedRoPEFuncBackward": "RoPE_bwd", "CrossEntropyFunctionBackward": "CrossEntropy_bwd", - # MoE auxiliary ops (sorting / topk) + # MoE auxiliary ops. "aiter::moe_sorting_fwd": "MoE_aux", "aiter::moe_sorting_opus_fwd": "MoE_aux", "aiter::moe_align_block_size": "MoE_aux", @@ -92,37 +73,34 @@ "aiter::biased_grouped_topk_hip": "MoE_aux", "aiter::grouped_topk": "MoE_aux", "aiter::moe_fused_gate": "MoE_aux", - # InferenceAttention extras (KV-cache writes) + # InferenceAttention extras (KV-cache writes). "_C_cache_ops::reshape_and_cache_flash": "InferenceAttention", "_C_cache_ops::concat_and_cache_mla": "InferenceAttention", } -# --------------------------------------------------------------------------- -# Patterns evaluated only when no exact registry match is found, in declared -# order. ``re.match`` is implicitly anchored at the start of the string. -# --------------------------------------------------------------------------- -OP_CATEGORY_PATTERNS: list[tuple[re.Pattern, str]] = [ - (re.compile(r"triton"), "triton"), - (re.compile(r"record_param_comms"), "record_param_comms"), +OP_CATEGORY_PATTERNS: List[Tuple[Pattern, str]] = [ + (re.compile(r"^triton"), "triton"), + (re.compile(r"^record_param_comms"), "record_param_comms"), ] -# --------------------------------------------------------------------------- -# Kernel-name fallback rules -# -# Applied only to GPU kernel names starting with ``void at::native``. The -# first substring match wins. Mirrors the kernel-name probe at the bottom -# of the legacy ``categorize_torch_op``. -# --------------------------------------------------------------------------- _KERNEL_NAME_PREFIX = "void at::native" -_KERNEL_NAME_FALLBACK_RULES: tuple[tuple[str, str], ...] = ( +_KERNEL_NAME_FALLBACK_RULES: Tuple[Tuple[str, str], ...] = ( ("elementwise", "elementwise"), ("reduce", "reduce"), ("multi_tensor_apply", "multi_tensor_apply"), ) +def _append_unique(target: List[str], names: Iterable[str]) -> None: + existing = set(target) + for name in names: + if name not in existing: + target.append(name) + existing.add(name) + + def _kernel_name_fallback(row) -> Optional[str]: kernel_details = row.get("kernel_details") if not kernel_details: @@ -136,33 +114,33 @@ def _kernel_name_fallback(row) -> Optional[str]: return None -def _resolve_base_category(op_name: str, base_category: str) -> str: - """Map ``(op_name, base-class category)`` to the final category emitted by v1. +def is_backward_op(op_name: str) -> bool: + return ( + op_name.endswith("_backward") + or op_name.endswith("Backward") + or op_name.endswith("_bwd") + or op_name in SDPA_BWD_OPS + ) - The legacy chain applied a handful of fwd/bwd splits and category renames - that aren't expressible by walking the perf-model class hierarchy alone. - This helper centralises those rules. - """ + +def resolve_base_category(op_name: str, base_category: str) -> str: + """Return the final category for an op in a base/sheet category.""" if base_category == "SDPA": - if op_name.endswith("_backward") or op_name in _LEGACY_SDPA_BWD_NAMES: - return "SDPA_bwd" - return "SDPA_fwd" + return "SDPA_bwd" if is_backward_op(op_name) else "SDPA_fwd" if base_category == "Normalization": - if op_name.endswith("_backward") or op_name.endswith("Backward"): - return "NORM_bwd" - return "NORM_fwd" + return "NORM_bwd" if is_backward_op(op_name) else "NORM_fwd" if base_category == "CONV": - if op_name.endswith("_backward") or op_name.endswith("Backward"): - return "CONV_bwd" - return "CONV_fwd" + return "CONV_bwd" if is_backward_op(op_name) else "CONV_fwd" + if base_category == "GroupedGEMM": + return "GroupedGEMM_bwd" if is_backward_op(op_name) else "GroupedGEMM_fwd" if base_category == "SSM": - return "SSM_fwd" + return "SSM_bwd" if is_backward_op(op_name) else "SSM_fwd" if base_category == "MoE_comm": - return "MoE_comm_fwd" + return "MoE_comm_bwd" if is_backward_op(op_name) else "MoE_comm_fwd" if base_category == "RoPE": - return "RoPE_fwd" + return "RoPE_bwd" if is_backward_op(op_name) else "RoPE_fwd" if base_category == "CrossEntropy": - return "CrossEntropy_fwd" + return "CrossEntropy_bwd" if is_backward_op(op_name) else "CrossEntropy_fwd" if base_category in ("BinaryElementwise", "UnaryElementwise"): return "elementwise" if base_category == "Reduce": @@ -170,74 +148,167 @@ def _resolve_base_category(op_name: str, base_category: str) -> str: return base_category +def sheet_category_from_final_category(category: str) -> str: + """Return the legacy sheet family for a final categorization label.""" + category_to_sheet = { + "CONV_fwd": "CONV", + "CONV_bwd": "CONV", + "SDPA_fwd": "SDPA", + "SDPA_bwd": "SDPA", + "NORM_fwd": "Normalization", + "NORM_bwd": "Normalization", + "GroupedGEMM_fwd": "GroupedGEMM", + "GroupedGEMM_bwd": "GroupedGEMM", + "SSM_fwd": "SSM", + "SSM_bwd": "SSM", + "MoE_comm_fwd": "MoE_comm", + "MoE_comm_bwd": "MoE_comm", + "RoPE_fwd": "RoPE", + "RoPE_bwd": "RoPE", + "CrossEntropy_fwd": "CrossEntropy", + "CrossEntropy_bwd": "CrossEntropy", + "elementwise": "UnaryElementwise", + "reduce": "Reduce", + } + return category_to_sheet.get(category, category) + + +def category_from_sheet_view( + op_name: str, dict_cat2names: Mapping[str, List[str]] +) -> Optional[str]: + """Compatibility fallback for callers that still mutate ``dict_cat2names``.""" + for sheet_category, names in dict_cat2names.items(): + if op_name in names: + return resolve_base_category(op_name, sheet_category) + return None + + def build_op_category_registry( op_to_perf_model_class_map: Mapping[str, type], dict_base_class2category: Mapping[type, str], - extras: Optional[Mapping[str, str]] = None, -) -> dict[str, str]: - """Construct the flat ``op_name -> category`` registry. - - The registry is built in two passes: - - 1. Walk every entry in ``op_to_perf_model_class_map``. Each perf-model - class contributes exactly one ``(op_name, category)`` pair, derived - from its single base class via ``dict_base_class2category`` and then - passed through :func:`_resolve_base_category` to apply fwd/bwd splits. - - 2. Apply ``extras`` as overrides. Use this for category-only ops - (no perf model) and for ops where the legacy chain disagrees with - the auto-derived category. - - Parameters - ---------- - op_to_perf_model_class_map - Mapping from op name to perf-model class. - dict_base_class2category - Mapping from perf-model base class to category label. - extras - Optional ``op_name -> category`` overrides. - - Returns - ------- - dict[str, str] - ``op_name -> final_category`` lookup table. - """ - registry: dict[str, str] = {} - + overrides: Optional[Mapping[str, str]] = None, +) -> Dict[str, str]: + """Construct the flat ``op_name -> final category`` registry.""" + registry: Dict[str, str] = {} for op_name, perf_model_class in op_to_perf_model_class_map.items(): base_classes = perf_model_class.__bases__ if len(base_classes) != 1: - continue + raise ValueError( + f"op_name: {op_name}, perf_model_class: {perf_model_class}, " + f"base_classes: {base_classes}" + ) base_category = dict_base_class2category.get(base_classes[0]) if base_category is None: - continue - registry[op_name] = _resolve_base_category(op_name, base_category) + raise ValueError( + f"op_name: {op_name}, perf_model_class: {perf_model_class}, " + f"base_class: {base_classes[0]}" + ) + registry[op_name] = resolve_base_category(op_name, base_category) - if extras: - registry.update(extras) + if overrides: + registry.update(overrides) return registry -def categorize_torch_op_v2( +def build_dict_cat2names( + op_to_perf_model_class_map: Mapping[str, type], + dict_base_class2category: Mapping[type, str], +) -> DefaultDict[str, List[str]]: + """Build the legacy ``category -> op names`` view used for report sheets.""" + dict_cat2names = defaultdict(list) # type: DefaultDict[str, List[str]] + for op_name, perf_model_class in op_to_perf_model_class_map.items(): + base_classes = perf_model_class.__bases__ + if len(base_classes) != 1: + raise ValueError( + f"op_name: {op_name}, perf_model_class: {perf_model_class}, " + f"base_classes: {base_classes}" + ) + base_category = dict_base_class2category.get(base_classes[0]) + if base_category is None: + raise ValueError( + f"op_name: {op_name}, perf_model_class: {perf_model_class}, " + f"base_class: {base_classes[0]}" + ) + dict_cat2names[base_category].append(op_name) + return dict_cat2names + + +def register_perf_model_categories( + perf_model_extension: Mapping[str, type], + dict_base_class2category: Mapping[type, str], + registry: MutableMapping[str, str], + dict_cat2names: MutableMapping[str, List[str]], +) -> None: + """Register categories for extension-provided perf models.""" + for op_name, perf_model_class in perf_model_extension.items(): + base_classes = perf_model_class.__bases__ + if len(base_classes) != 1: + raise ValueError( + f"op_name: {op_name}, perf_model_class: {perf_model_class}, " + f"base_classes: {base_classes}" + ) + sheet_category = dict_base_class2category.get(base_classes[0]) + if sheet_category is None: + raise ValueError( + f"op_name: {op_name}, perf_model_class: {perf_model_class}, " + f"base_class: {base_classes[0]}" + ) + registry[op_name] = resolve_base_category(op_name, sheet_category) + if sheet_category not in dict_cat2names: + dict_cat2names[sheet_category] = [] + _append_unique(dict_cat2names[sheet_category], [op_name]) + + +def register_op_categories( + op_category_extension: Mapping[str, str], + registry: MutableMapping[str, str], + dict_cat2names: Optional[MutableMapping[str, List[str]]] = None, +) -> None: + """Register explicit category-only op labels.""" + registry.update(op_category_extension) + if dict_cat2names is None: + return + for op_name, category in op_category_extension.items(): + sheet_category = sheet_category_from_final_category(category) + if sheet_category not in dict_cat2names: + dict_cat2names[sheet_category] = [] + _append_unique(dict_cat2names[sheet_category], [op_name]) + + +def register_dict_cat2names_extension( + dict_cat2names_extension: Mapping[str, List[str]], + registry: MutableMapping[str, str], + dict_cat2names: MutableMapping[str, List[str]], +) -> None: + """Support the older ``dict_cat2names_extension`` extension contract.""" + for sheet_category, names in dict_cat2names_extension.items(): + if not isinstance(names, list): + raise ValueError(f"Expected names to be a list, got {type(names)}") + if sheet_category not in dict_cat2names: + dict_cat2names[sheet_category] = [] + _append_unique(dict_cat2names[sheet_category], names) + for op_name in names: + registry[op_name] = resolve_base_category(op_name, sheet_category) + + +def categorize_torch_op_from_registry( row, registry: Mapping[str, str], - patterns: Iterable[tuple[re.Pattern, str]] = OP_CATEGORY_PATTERNS, + dict_cat2names: Optional[Mapping[str, List[str]]] = None, + patterns: Iterable[Tuple[Pattern, str]] = OP_CATEGORY_PATTERNS, ) -> str: - """Return the category for ``row`` using the flat registry plus patterns. - - Resolution order: - - 1. Exact name match in ``registry``. - 2. First pattern in ``patterns`` whose regex matches ``row["name"]``. - 3. Kernel-name fallback for ``void at::native`` GPU kernels. - 4. ``"other"``. - """ + """Return the category for ``row`` using explicit registry data.""" name = row["name"] - cat = registry.get(name) - if cat is not None: - return cat + category = registry.get(name) + if category is not None: + return category + + if dict_cat2names is not None: + category = category_from_sheet_view(name, dict_cat2names) + if category is not None: + return category for pattern, category in patterns: if pattern.match(name): diff --git a/TraceLens/PerfModel/torch_op_mapping.py b/TraceLens/PerfModel/torch_op_mapping.py index e4aa33b84..d063af6d1 100644 --- a/TraceLens/PerfModel/torch_op_mapping.py +++ b/TraceLens/PerfModel/torch_op_mapping.py @@ -5,12 +5,12 @@ ############################################################################### from . import perf_model -from collections import defaultdict from .extensions import get_pseudo_op_mappings, get_pseudo_op_categories from .op_categories import ( - LEGACY_CATEGORIZE_EXTRAS, + OP_CATEGORY_OVERRIDES, + build_dict_cat2names, build_op_category_registry, - categorize_torch_op_v2 as _categorize_v2, + categorize_torch_op_from_registry, ) op_to_perf_model_class_map = { @@ -202,7 +202,7 @@ dict_base_class2category = { perf_model.GEMM: "GEMM", - perf_model.GroupedGemm: "GEMM", + perf_model.GroupedGemm: "GroupedGEMM", perf_model.CONV: "CONV", perf_model.SDPA: "SDPA", perf_model.UnaryElementwise: "UnaryElementwise", @@ -219,19 +219,16 @@ # Add pseudo-op extension categories dict_base_class2category.update(get_pseudo_op_categories()) -dict_cat2names = defaultdict(list) -for op_name, perf_model_class in op_to_perf_model_class_map.items(): - base_classes = perf_model_class.__bases__ - assert ( - len(base_classes) == 1 - ), f"op_name: {op_name}, perf_model_class: {perf_model_class}, base_classes: {base_classes}" - base_class = base_classes[0] - cat = dict_base_class2category.get(base_class) - if cat is None: - raise ValueError( - f"op_name: {op_name}, perf_model_class: {perf_model_class}, base_class: {base_classes}" - ) - dict_cat2names[cat].append(op_name) +dict_cat2names = build_dict_cat2names( + op_to_perf_model_class_map, + dict_base_class2category, +) + +OP_CATEGORY_REGISTRY = build_op_category_registry( + op_to_perf_model_class_map, + dict_base_class2category, + overrides=OP_CATEGORY_OVERRIDES, +) def categorize_torch_op(row): @@ -245,7 +242,8 @@ def categorize_torch_op(row): Returns: str: One of 'GEMM', 'CONV_fwd', 'CONV_bwd', 'NORM_fwd', 'NORM_bwd', - 'SDPA_fwd', 'SDPA_bwd', 'MoE_fused', 'MoE_unfused', + 'SDPA_fwd', 'SDPA_bwd', 'GroupedGEMM_fwd', 'GroupedGEMM_bwd', + 'MoE_fused', 'MoE_unfused', 'SSM_fwd', 'SSM_bwd', 'MoE_comm_fwd', 'MoE_comm_bwd', 'RoPE_fwd', 'RoPE_bwd', 'CrossEntropy_fwd', 'CrossEntropy_bwd', 'elementwise', 'triton', 'reduce', 'multi_tensor_apply', @@ -254,175 +252,8 @@ def categorize_torch_op(row): Note: Backward variants and auxiliary ops (TokenPermuteMaskMap, etc.) are categorization-only (timing without GFLOPS or TB/s). """ - - debug = False - if row["name"] in dict_cat2names["GEMM"]: - return "GEMM" - elif row["name"] in [ - "aten::convolution", - "aten::miopen_convolution", - "aten::cudnn_convolution", - "ConvBias_", - "ConvBiasReLU_", - ]: - return "CONV_fwd" - elif row["name"] in [ - "aten::convolution_backward", - "ConvBias_Backward", - "ConvBiasReLU_Backward", - ]: - return "CONV_bwd" - elif row["name"] in norm_ops.keys() or row["name"] in dict_cat2names.get( - "Normalization", [] - ): - if row["name"].endswith("_backward") or row["name"].endswith("Backward"): - return "NORM_bwd" - else: - return "NORM_fwd" - # SDPA ops: distinguish forward and backward - sdpa_bwd_names = [ - "FlashAttnFuncBackward", - "FusedAttnFuncBackward", - "flash_attn::_flash_attn_backward", - "flash_attn::_flash_attn_varlen_backward", - "aten::_scaled_dot_product_cudnn_attention_backward", - "aten::_scaled_dot_product_efficient_attention_backward", - "aten::_scaled_dot_product_flash_attention_backward", - "aiter::_flash_attn_backward", - "aiter::wrapper_fmha_v3_bwd", - "aiter::mha_bwd", - ] - if row["name"] in dict_cat2names["SDPA"]: - if row["name"].endswith("_backward") or row["name"] in sdpa_bwd_names: - return "SDPA_bwd" - else: - return "SDPA_fwd" - elif row["name"] in dict_cat2names.get("GroupedGEMM", []): - if row["name"].endswith("Backward"): - return "GroupedGEMM_bwd" - else: - return "GroupedGEMM_fwd" - elif row["name"] in dict_cat2names.get("MoE_fused", []): - return "MoE_fused" - elif row["name"] in dict_cat2names.get("MoE_unfused", []): - return "MoE_unfused" - elif row["name"] in dict_cat2names.get("SSM", []) or row["name"] in [ - "MambaSplitConv1dScanCombinedFn", - ]: - return "SSM_fwd" - elif row["name"] in [ - "MambaSplitConv1dScanCombinedFnBackward", - "DaoAILab::_causal_conv1d_bwd_cpp", - ]: - return "SSM_bwd" - elif row["name"] in dict_cat2names.get("MoE_comm", []) or row["name"] in [ - "TokenPermuteMaskMap", - "_OperationFuserAutogradFunction", - ]: - return "MoE_comm_fwd" - elif row["name"] in [ - "MoEDispatchBackward", - "MoECombineBackward", - "TokenPermuteMaskMapBackward", - "_OperationFuserAutogradFunctionBackward", - ]: - return "MoE_comm_bwd" - elif row["name"] in dict_cat2names.get("RoPE", []): - return "RoPE_fwd" - elif row["name"] in ["FusedRoPEFuncBackward"]: - return "RoPE_bwd" - elif row["name"] in dict_cat2names.get("CrossEntropy", []): - return "CrossEntropy_fwd" - elif row["name"] in ["CrossEntropyFunctionBackward"]: - return "CrossEntropy_bwd" - elif row["name"] in dict_cat2names.get("BinaryElementwise", []): - return "elementwise" - elif row["name"] in dict_cat2names.get("Reduce", []): - return "reduce" - elif row["name"].startswith("triton"): - return "triton" - elif row["name"] in [ - "aiter::moe_sorting_fwd", - "aiter::moe_sorting_opus_fwd", - "aiter::moe_align_block_size", - "_moe_C::moe_align_block_size", - "aiter::fused_moe_->_fused_dynamic_mxfp4_quant_moe_sort_kernel (Synthetic Op)", - ]: - return "MoE_aux" - elif row["name"] in [ - "aiter::moe_sum", - ]: - return "MoE_aux" - elif row["name"] in [ - "aiter::topk_softmax", - "aiter::topk_softmax_asm", - "aiter::topk_sigmoid", - "aiter::biased_grouped_topk_hip", - "aiter::grouped_topk", - "aiter::moe_fused_gate", - ]: - return "MoE_aux" - elif row["name"] in [ - "_C_cache_ops::reshape_and_cache_flash", - "_C_cache_ops::concat_and_cache_mla", - ]: - return "InferenceAttention" - elif row["name"].startswith("record_param_comms"): - return "record_param_comms" - elif row["name"] in dict_cat2names.get("MoE_fused", []): - return "MoE_fused" - elif row["name"] in dict_cat2names.get("MoE_unfused", []): - return "MoE_unfused" - elif row["name"] in dict_cat2names.get("InferenceAttention", []): - return "InferenceAttention" - elif row["name"] in dict_cat2names.get("RMSNorm", []): - return "RMSNorm" - elif row["name"] in dict_cat2names.get("CustomCollective", []): - return "CustomCollective" - elif row["name"] in dict_cat2names.get("GroupQuant", []): - return "GroupQuant" - elif row["name"] in dict_cat2names.get("BinaryElementwise", []): - return "elementwise" - elif row["name"] in dict_cat2names.get("UnaryElementwise", []): - return "elementwise" - elif row["name"] in dict_cat2names.get("Reduce", []): - return "reduce" - if "kernel_details" in row and len(row["kernel_details"]) > 0: - kernel_name = row["kernel_details"][0]["name"] - # else: - # raise ValueError( - # f"Row does not contain 'kernel_names' or 'kernel_details' with a valid name. Row: {row}" - # ) - if kernel_name.startswith("void at::native"): - if debug: - print("Found ATen native kernel:", kernel_name[:64]) - if "elementwise" in kernel_name: - return "elementwise" - elif "reduce" in kernel_name: - return "reduce" - elif "multi_tensor_apply" in kernel_name: - return "multi_tensor_apply" - # if none of the above cases match, return 'other' - return "other" - - -# --------------------------------------------------------------------------- -# Registry-based v2 categorizer (introduced in PR A; runs alongside the -# legacy ``categorize_torch_op`` chain above). The parity test -# ``tests/test_categorize_torch_op_parity.py`` asserts both produce the same -# output for every reachable name. PR B will swap v2 in and delete v1. -# --------------------------------------------------------------------------- -OP_CATEGORY_REGISTRY = build_op_category_registry( - op_to_perf_model_class_map, - dict_base_class2category, - extras=LEGACY_CATEGORIZE_EXTRAS, -) - - -def categorize_torch_op_v2(row): - """Registry-based replacement for :func:`categorize_torch_op`. - - Equivalent in behavior to the legacy if/elif chain (verified by the - parity test). See ``op_categories.py`` for the design rationale. - """ - return _categorize_v2(row, OP_CATEGORY_REGISTRY) + return categorize_torch_op_from_registry( + row, + OP_CATEGORY_REGISTRY, + dict_cat2names=dict_cat2names, + ) diff --git a/TraceLens/Reporting/generate_perf_report_pytorch.py b/TraceLens/Reporting/generate_perf_report_pytorch.py index 37c018f03..fbc46327a 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch.py @@ -119,6 +119,16 @@ def apply_extension(perf_analyzer, extension_path): extension_path = os.path.abspath(extension_path) extension_name = os.path.splitext(os.path.basename(extension_path))[0] + from TraceLens.PerfModel.op_categories import ( + register_dict_cat2names_extension, + register_op_categories, + register_perf_model_categories, + ) + from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, + dict_base_class2category, + ) + spec = importlib.util.spec_from_file_location(extension_name, extension_path) extension = importlib.util.module_from_spec(spec) spec.loader.exec_module(extension) @@ -137,6 +147,24 @@ def apply_extension(perf_analyzer, extension_path): f"Expected perf_model_extension to be a dict, got {type(perf_model_extension)}" ) perf_analyzer.op_to_perf_model_class_map.update(perf_model_extension) + register_perf_model_categories( + perf_model_extension, + dict_base_class2category, + OP_CATEGORY_REGISTRY, + perf_analyzer.dict_cat2names, + ) + if hasattr(extension, "op_category_extension"): + print(f"Applying op category extension from {extension_path}") + op_category_extension = getattr(extension, "op_category_extension") + if not isinstance(op_category_extension, dict): + raise ValueError( + f"Expected op_category_extension to be a dict, got {type(op_category_extension)}" + ) + register_op_categories( + op_category_extension, + OP_CATEGORY_REGISTRY, + perf_analyzer.dict_cat2names, + ) if hasattr(extension, "dict_cat2names_extension"): print(f"Updating dict_cat2names with extension from {extension_path}") if not isinstance(extension.dict_cat2names_extension, dict): @@ -144,13 +172,11 @@ def apply_extension(perf_analyzer, extension_path): f"Expected dict_cat2names_extension to be a dict, got {type(extension.dict_cat2names_extension)}" ) - # defaultdict(, - for cat, names in extension.dict_cat2names_extension.items(): - if cat not in perf_analyzer.dict_cat2names: - perf_analyzer.dict_cat2names[cat] = [] - if not isinstance(names, list): - raise ValueError(f"Expected names to be a list, got {type(names)}") - perf_analyzer.dict_cat2names[cat].extend(names) + register_dict_cat2names_extension( + extension.dict_cat2names_extension, + OP_CATEGORY_REGISTRY, + perf_analyzer.dict_cat2names, + ) def trunc_kernel_details(row, kernel_detail_col, trunc_length=64): diff --git a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py index 40ce4d58c..8dc125e1d 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py @@ -400,6 +400,16 @@ def apply_extension(perf_analyzer, extension_path): extension_path = os.path.abspath(extension_path) extension_name = os.path.splitext(os.path.basename(extension_path))[0] + from TraceLens.PerfModel.op_categories import ( + register_dict_cat2names_extension, + register_op_categories, + register_perf_model_categories, + ) + from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, + dict_base_class2category, + ) + spec = importlib.util.spec_from_file_location(extension_name, extension_path) extension = importlib.util.module_from_spec(spec) spec.loader.exec_module(extension) @@ -418,6 +428,24 @@ def apply_extension(perf_analyzer, extension_path): f"Expected perf_model_extension to be a dict, got {type(perf_model_extension)}" ) perf_analyzer.op_to_perf_model_class_map.update(perf_model_extension) + register_perf_model_categories( + perf_model_extension, + dict_base_class2category, + OP_CATEGORY_REGISTRY, + perf_analyzer.dict_cat2names, + ) + if hasattr(extension, "op_category_extension"): + print(f"Applying op category extension from {extension_path}") + op_category_extension = getattr(extension, "op_category_extension") + if not isinstance(op_category_extension, dict): + raise ValueError( + f"Expected op_category_extension to be a dict, got {type(op_category_extension)}" + ) + register_op_categories( + op_category_extension, + OP_CATEGORY_REGISTRY, + perf_analyzer.dict_cat2names, + ) if hasattr(extension, "dict_cat2names_extension"): print(f"Updating dict_cat2names with extension from {extension_path}") if not isinstance(extension.dict_cat2names_extension, dict): @@ -425,13 +453,11 @@ def apply_extension(perf_analyzer, extension_path): f"Expected dict_cat2names_extension to be a dict, got {type(extension.dict_cat2names_extension)}" ) - # defaultdict(, - for cat, names in extension.dict_cat2names_extension.items(): - if cat not in perf_analyzer.dict_cat2names: - perf_analyzer.dict_cat2names[cat] = [] - if not isinstance(names, list): - raise ValueError(f"Expected names to be a list, got {type(names)}") - perf_analyzer.dict_cat2names[cat].extend(names) + register_dict_cat2names_extension( + extension.dict_cat2names_extension, + OP_CATEGORY_REGISTRY, + perf_analyzer.dict_cat2names, + ) def trunc_kernel_details(row, kernel_detail_col, trunc_length=64): diff --git a/tests/test_categorize_torch_op_parity.py b/tests/test_categorize_torch_op_parity.py deleted file mode 100644 index 6122a0004..000000000 --- a/tests/test_categorize_torch_op_parity.py +++ /dev/null @@ -1,197 +0,0 @@ -############################################################################### -# Copyright (c) 2024 - 2025 Advanced Micro Devices, Inc. All rights reserved. -# -# See LICENSE for license information. -############################################################################### - -""" -Parity test for :func:`categorize_torch_op_v2` (registry-based) against -:func:`categorize_torch_op` (legacy if/elif chain). - -PR A introduces v2 alongside v1 without changing v1's behavior. The single -acceptance criterion is that v2 produces the same output as v1 for every -reachable input. PR B will delete v1; that PR's safety relies on this test. -""" - -import pytest - -from TraceLens.PerfModel.torch_op_mapping import ( - OP_CATEGORY_REGISTRY, - categorize_torch_op, - categorize_torch_op_v2, - op_to_perf_model_class_map, -) - - -def _row(name, kernel_details=None): - return {"name": name, "kernel_details": kernel_details or []} - - -# --------------------------------------------------------------------------- -# Coverage set 1: every name reachable through perf-model registration -# (op_to_perf_model_class_map, including extension-contributed entries). -# --------------------------------------------------------------------------- -PERF_MODEL_NAMES = sorted(op_to_perf_model_class_map.keys()) - - -# --------------------------------------------------------------------------- -# Coverage set 2: names hardcoded inside the legacy categorize_torch_op -# if/elif chain. These are the names PR A had to thread into -# LEGACY_CATEGORIZE_EXTRAS to preserve parity. -# --------------------------------------------------------------------------- -LEGACY_HARDCODED_NAMES = [ - # CONV_fwd - "aten::convolution", - "aten::miopen_convolution", - "aten::cudnn_convolution", - "ConvBias_", - "ConvBiasReLU_", - # CONV_bwd - "aten::convolution_backward", - "ConvBias_Backward", - "ConvBiasReLU_Backward", - # SDPA backward names (some perf-modeled, some unreachable in v1) - "FlashAttnFuncBackward", - "FusedAttnFuncBackward", - "flash_attn::_flash_attn_backward", - "flash_attn::_flash_attn_varlen_backward", - "aten::_scaled_dot_product_cudnn_attention_backward", - "aten::_scaled_dot_product_efficient_attention_backward", - "aten::_scaled_dot_product_flash_attention_backward", - "aiter::_flash_attn_backward", - "aiter::wrapper_fmha_v3_bwd", - "aiter::mha_bwd", - # SSM - "MambaSplitConv1dScanCombinedFn", - "MambaSplitConv1dScanCombinedFnBackward", - "DaoAILab::_causal_conv1d_bwd_cpp", - # MoE_comm - "TokenPermuteMaskMap", - "_OperationFuserAutogradFunction", - "MoEDispatchBackward", - "MoECombineBackward", - "TokenPermuteMaskMapBackward", - "_OperationFuserAutogradFunctionBackward", - # RoPE / CrossEntropy backward - "FusedRoPEFuncBackward", - "CrossEntropyFunctionBackward", - # MoE auxiliary - "aiter::moe_sorting_fwd", - "aiter::moe_sorting_opus_fwd", - "aiter::moe_align_block_size", - "_moe_C::moe_align_block_size", - "aiter::fused_moe_->_fused_dynamic_mxfp4_quant_moe_sort_kernel (Synthetic Op)", - "aiter::moe_sum", - "aiter::topk_softmax", - "aiter::topk_softmax_asm", - "aiter::topk_sigmoid", - "aiter::biased_grouped_topk_hip", - "aiter::grouped_topk", - "aiter::moe_fused_gate", - # InferenceAttention extras - "_C_cache_ops::reshape_and_cache_flash", - "_C_cache_ops::concat_and_cache_mla", -] - - -# --------------------------------------------------------------------------- -# Coverage set 3: names matched by patterns (triton, record_param_comms). -# --------------------------------------------------------------------------- -PATTERN_NAMES = [ - "triton", - "triton_softmax_kernel", - "triton_per_token_quant_fp8", - "record_param_comms", - "record_param_comms_alltoall_base", -] - - -# --------------------------------------------------------------------------- -# Coverage set 4: unknown names that should resolve to "other" (or to -# a kernel-name fallback). -# --------------------------------------------------------------------------- -OTHER_NAMES = [ - "completely::unknown::op", - "no_match_at_all", - "Some::Random::Op", -] - - -# --------------------------------------------------------------------------- -# Coverage set 5: kernel-name fallback inputs. -# --------------------------------------------------------------------------- -KERNEL_FALLBACK_CASES = [ - # (kernel_name, expected_category) - ("void at::native::elementwise_kernel<8, 4>(...)", "elementwise"), - ("void at::native::reduce_kernel<512, 4>(...)", "reduce"), - ("void at::native::multi_tensor_apply_kernel<...>(...)", "multi_tensor_apply"), - # Native kernel that matches no needle -> "other" - ("void at::native::unrelated_kernel(...)", "other"), - # Non-native kernel -> "other" - ("some_user_kernel(...)", "other"), -] - - -# --------------------------------------------------------------------------- -# Parametrized parity assertions -# --------------------------------------------------------------------------- -@pytest.mark.parametrize("name", PERF_MODEL_NAMES) -def test_parity_perf_model_names(name): - row = _row(name) - assert categorize_torch_op_v2(row) == categorize_torch_op(row) - - -@pytest.mark.parametrize("name", LEGACY_HARDCODED_NAMES) -def test_parity_legacy_hardcoded_names(name): - row = _row(name) - assert categorize_torch_op_v2(row) == categorize_torch_op(row) - - -@pytest.mark.parametrize("name", PATTERN_NAMES) -def test_parity_pattern_names(name): - row = _row(name) - assert categorize_torch_op_v2(row) == categorize_torch_op(row) - - -@pytest.mark.parametrize("name", OTHER_NAMES) -def test_parity_other_names(name): - row = _row(name) - assert categorize_torch_op_v2(row) == categorize_torch_op(row) - - -@pytest.mark.parametrize("kernel_name,expected", KERNEL_FALLBACK_CASES) -def test_parity_kernel_fallback(kernel_name, expected): - row = _row("unknown_op_name", kernel_details=[{"name": kernel_name}]) - v1 = categorize_torch_op(row) - v2 = categorize_torch_op_v2(row) - assert v1 == v2 == expected, f"v1={v1!r}, v2={v2!r}, expected={expected!r}" - - -# --------------------------------------------------------------------------- -# Sanity checks on the registry shape -# --------------------------------------------------------------------------- -def test_registry_built_at_import(): - assert isinstance(OP_CATEGORY_REGISTRY, dict) - assert len(OP_CATEGORY_REGISTRY) > 0 - - -def test_registry_contains_expected_entries(): - # Auto-derived from a perf-model base class - assert OP_CATEGORY_REGISTRY["aten::mm"] == "GEMM" - # Auto-derived NORM split - assert OP_CATEGORY_REGISTRY["aten::layer_norm"] == "NORM_fwd" - assert OP_CATEGORY_REGISTRY["aten::layer_norm_backward"] == "NORM_bwd" - # Auto-derived SDPA fwd vs explicit-list bwd - assert OP_CATEGORY_REGISTRY["FlashAttnFunc"] == "SDPA_fwd" - assert OP_CATEGORY_REGISTRY["aiter::wrapper_fmha_v3_bwd"] == "SDPA_bwd" - # Extras-only entries - assert OP_CATEGORY_REGISTRY["aten::miopen_convolution"] == "CONV_fwd" - assert OP_CATEGORY_REGISTRY["DaoAILab::_causal_conv1d_bwd_cpp"] == "SSM_bwd" - assert OP_CATEGORY_REGISTRY["aiter::topk_softmax"] == "MoE_aux" - - -def test_registry_covers_every_perf_model_op(): - # The builder must produce a category for every perf-modeled op (i.e. no - # silent drops). PR B relies on this invariant. - missing = [n for n in op_to_perf_model_class_map if n not in OP_CATEGORY_REGISTRY] - assert missing == [], f"perf-modeled ops missing from registry: {missing}" diff --git a/tests/test_primus_op_categorization.py b/tests/test_primus_op_categorization.py index df74074bb..645330539 100644 --- a/tests/test_primus_op_categorization.py +++ b/tests/test_primus_op_categorization.py @@ -43,17 +43,17 @@ def test_ck_grouped_gemm_variable_k_mapped(): assert op_to_perf_model_class_map[op] is primus_turbo_grouped_gemm_variable_k -def test_ck_grouped_gemm_categorizes_as_gemm(): +def test_ck_grouped_gemm_categorizes_as_grouped_gemm(): row = {"name": "primus_turbo_cpp_extension::ck_grouped_gemm", "kernel_details": []} - assert categorize_torch_op(row) == "GEMM" + assert categorize_torch_op(row) == "GroupedGEMM_fwd" -def test_ck_grouped_gemm_variable_k_categorizes_as_gemm(): +def test_ck_grouped_gemm_variable_k_categorizes_as_grouped_gemm(): row = { "name": "primus_turbo_cpp_extension::ck_grouped_gemm_variable_k", "kernel_details": [], } - assert categorize_torch_op(row) == "GEMM" + assert categorize_torch_op(row) == "GroupedGEMM_fwd" def test_ck_grouped_gemm_flops(): diff --git a/tests/test_primus_turbo_grouped_gemm.py b/tests/test_primus_turbo_grouped_gemm.py index a62733c93..50720967c 100644 --- a/tests/test_primus_turbo_grouped_gemm.py +++ b/tests/test_primus_turbo_grouped_gemm.py @@ -27,7 +27,7 @@ def test_fixed_k_ops_map_to_primus_turbo_grouped_gemm(): for op in fixed_k_ops: assert op_to_perf_model_class_map[op] is primus_turbo_grouped_gemm, op row = {"name": op, "kernel_details": []} - assert categorize_torch_op(row) == "GEMM", op + assert categorize_torch_op(row) == "GroupedGEMM_fwd", op def test_variable_k_ops_map_to_primus_turbo_grouped_gemm_variable_k(): @@ -41,7 +41,7 @@ def test_variable_k_ops_map_to_primus_turbo_grouped_gemm_variable_k(): op_to_perf_model_class_map[op] is primus_turbo_grouped_gemm_variable_k ), op row = {"name": op, "kernel_details": []} - assert categorize_torch_op(row) == "GEMM", op + assert categorize_torch_op(row) == "GroupedGEMM_fwd", op # --------------------------------------------------------------------------- diff --git a/tests/test_torch_op_categorization_registry.py b/tests/test_torch_op_categorization_registry.py new file mode 100644 index 000000000..24ec031ac --- /dev/null +++ b/tests/test_torch_op_categorization_registry.py @@ -0,0 +1,138 @@ +############################################################################### +# Copyright (c) 2024 - 2025 Advanced Micro Devices, Inc. All rights reserved. +# +# See LICENSE for license information. +############################################################################### + +""" +Tests for the registry-based torch op categorizer. +""" + +import pytest + +from TraceLens.PerfModel.op_categories import ( + category_from_sheet_view, + register_dict_cat2names_extension, + register_op_categories, +) +from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, + categorize_torch_op, + dict_cat2names, + op_to_perf_model_class_map, +) + + +def _row(name, kernel_details=None): + return {"name": name, "kernel_details": kernel_details or []} + + +@pytest.mark.parametrize( + "name,expected", + [ + ("aten::mm", "GEMM"), + ("aten::addmm", "GEMM"), + ("aten::convolution", "CONV_fwd"), + ("aten::miopen_convolution", "CONV_fwd"), + ("aten::cudnn_convolution", "CONV_fwd"), + ("aten::convolution_backward", "CONV_bwd"), + ("ConvBias_Backward", "CONV_bwd"), + ("aten::layer_norm", "NORM_fwd"), + ("aten::native_layer_norm_backward", "NORM_bwd"), + ("FlashAttnFunc", "SDPA_fwd"), + ("FlashAttnFuncBackward", "SDPA_bwd"), + ("FusedAttnFuncBackward", "SDPA_bwd"), + ("flash_attn::_flash_attn_backward", "SDPA_bwd"), + ("aten::_scaled_dot_product_flash_attention_backward", "SDPA_bwd"), + ("primus_turbo::grouped_gemm", "GroupedGEMM_fwd"), + ("primus_turbo_cpp_extension::ck_grouped_gemm", "GroupedGEMM_fwd"), + ("MambaSplitConv1dScanCombinedFn", "SSM_fwd"), + ("MambaSplitConv1dScanCombinedFnBackward", "SSM_bwd"), + ("DaoAILab::_causal_conv1d_bwd_cpp", "SSM_bwd"), + ("MoEDispatch", "MoE_comm_fwd"), + ("MoEDispatchBackward", "MoE_comm_bwd"), + ("TokenPermuteMaskMap", "MoE_comm_fwd"), + ("TokenPermuteMaskMapBackward", "MoE_comm_bwd"), + ("FusedRoPEFunc", "RoPE_fwd"), + ("FusedRoPEFuncBackward", "RoPE_bwd"), + ("CrossEntropyFunction", "CrossEntropy_fwd"), + ("CrossEntropyFunctionBackward", "CrossEntropy_bwd"), + ("aiter::topk_softmax", "MoE_aux"), + ("_C_cache_ops::reshape_and_cache_flash", "InferenceAttention"), + ("aiter::silu_and_mul", "elementwise"), + ("aten::sum", "reduce"), + ("triton_per_token_quant_fp8", "triton"), + ("record_param_comms_alltoall_base", "record_param_comms"), + ("completely::unknown::op", "other"), + ], +) +def test_categorize_torch_op_expected_categories(name, expected): + assert categorize_torch_op(_row(name)) == expected + + +@pytest.mark.parametrize( + "kernel_name,expected", + [ + ("void at::native::elementwise_kernel<8, 4>(...)", "elementwise"), + ("void at::native::reduce_kernel<512, 4>(...)", "reduce"), + ( + "void at::native::multi_tensor_apply_kernel<...>(...)", + "multi_tensor_apply", + ), + ("void at::native::unrelated_kernel(...)", "other"), + ("some_user_kernel(...)", "other"), + ], +) +def test_kernel_name_fallback(kernel_name, expected): + row = _row("unknown_op_name", kernel_details=[{"name": kernel_name}]) + assert categorize_torch_op(row) == expected + + +def test_registry_covers_every_perf_model_op(): + missing = [name for name in op_to_perf_model_class_map if name not in OP_CATEGORY_REGISTRY] + assert missing == [] + + +def test_dict_cat2names_keeps_sheet_compatibility_view(): + assert "aten::mm" in dict_cat2names["GEMM"] + assert "primus_turbo::grouped_gemm" in dict_cat2names["GroupedGEMM"] + assert "aten::convolution" in dict_cat2names["CONV"] + assert "aten::convolution_backward" in dict_cat2names["CONV"] + assert "FlashAttnFuncBackward" not in dict_cat2names["SDPA"] + + +def test_dict_cat2names_dynamic_fallback_for_legacy_callers(): + local_sheet_view = {"SDPA": ["MyCustomAttentionBackward"]} + assert ( + category_from_sheet_view("MyCustomAttentionBackward", local_sheet_view) + == "SDPA_bwd" + ) + + +def test_register_op_category_extension_updates_registry_and_sheet_view(): + registry = {} + sheet_view = {} + + register_op_categories( + {"MyCategoryOnlyBackward": "SDPA_bwd"}, + registry, + sheet_view, + ) + + assert registry["MyCategoryOnlyBackward"] == "SDPA_bwd" + assert sheet_view["SDPA"] == ["MyCategoryOnlyBackward"] + + +def test_register_dict_cat2names_extension_updates_registry_and_sheet_view(): + registry = {} + sheet_view = {} + + register_dict_cat2names_extension( + {"GroupedGEMM": ["GroupedGemm", "GroupedGemmBackward"]}, + registry, + sheet_view, + ) + + assert registry["GroupedGemm"] == "GroupedGEMM_fwd" + assert registry["GroupedGemmBackward"] == "GroupedGEMM_bwd" + assert sheet_view["GroupedGEMM"] == ["GroupedGemm", "GroupedGemmBackward"] From 0aa179f006633beaa4dd17ba57c193b4f88e24d7 Mon Sep 17 00:00:00 2001 From: Jassani Date: Thu, 7 May 2026 18:41:01 -0400 Subject: [PATCH 03/10] PerfModel: declare op categories on model classes Move perf-modeled op categorization onto category/bwd_category metadata so extension models define their own report category. Keep category-only ops separate from legacy sheet membership and expose linked-forward backward metrics in the unified perf output. Co-authored-by: Cursor --- .../attention_perf_model_extensions.py | 3 + ...ustom_collectives_perf_model_extensions.py | 3 + .../extensions/moe_perf_model_extensions.py | 9 + .../extensions/perf_model_extensions.py | 3 + .../rmsnorm_perf_model_extensions.py | 21 ++ TraceLens/PerfModel/op_categories.py | 187 +++++++----------- TraceLens/PerfModel/perf_model.py | 59 ++++++ TraceLens/PerfModel/torch_op_mapping.py | 34 +--- .../Reporting/generate_perf_report_pytorch.py | 22 +-- .../generate_perf_report_pytorch_inference.py | 22 +-- TraceLens/TreePerf/tree_perf.py | 34 +++- docs/generate_perf_report.md | 2 +- examples/example_megatron_extension.py | 19 +- tests/test_pseudo_ops_extension.py | 68 +++---- .../test_torch_op_categorization_registry.py | 22 +-- 15 files changed, 258 insertions(+), 250 deletions(-) diff --git a/TraceLens/PerfModel/extensions/attention_perf_model_extensions.py b/TraceLens/PerfModel/extensions/attention_perf_model_extensions.py index b73ed6173..fe765ed53 100644 --- a/TraceLens/PerfModel/extensions/attention_perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/attention_perf_model_extensions.py @@ -30,6 +30,9 @@ class InferenceAttention: ``params.get("_no_perf")`` is true. """ + category = "InferenceAttention" + bwd_category = None + REQUIRED_PARAM_KEYS = ( "B", "N_Q", diff --git a/TraceLens/PerfModel/extensions/custom_collectives_perf_model_extensions.py b/TraceLens/PerfModel/extensions/custom_collectives_perf_model_extensions.py index adcbd1234..3aea2a303 100644 --- a/TraceLens/PerfModel/extensions/custom_collectives_perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/custom_collectives_perf_model_extensions.py @@ -14,6 +14,9 @@ class CustomCollective: + category = "CustomCollective" + bwd_category = None + pass diff --git a/TraceLens/PerfModel/extensions/moe_perf_model_extensions.py b/TraceLens/PerfModel/extensions/moe_perf_model_extensions.py index 638212bc8..0a4295b56 100644 --- a/TraceLens/PerfModel/extensions/moe_perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/moe_perf_model_extensions.py @@ -44,6 +44,9 @@ class FusedMoE: activation, and down projection) into a single kernel launch. """ + category = "MoE_fused" + bwd_category = None + def __init__(self, event, arch=None, python_path=None): self.event = event self.arch = arch @@ -386,6 +389,9 @@ class UnfusedMoE_Up: - Optionally gated (e.g., SwiGLU): both up and gate projections """ + category = "MoE_unfused" + bwd_category = None + @staticmethod def flops_func(num_tokens, hidden_dim, inter_dim, topk, gated): """ @@ -475,6 +481,9 @@ class UnfusedMoE_Down: Child classes implement get_param_details() to extract parameters from events. """ + category = "MoE_unfused" + bwd_category = None + @staticmethod def flops_func(num_tokens, hidden_dim, inter_dim, topk): """ diff --git a/TraceLens/PerfModel/extensions/perf_model_extensions.py b/TraceLens/PerfModel/extensions/perf_model_extensions.py index 6c42e4895..cfa4894b4 100644 --- a/TraceLens/PerfModel/extensions/perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/perf_model_extensions.py @@ -375,6 +375,9 @@ class GroupQuant(BinaryElementwise): Performance model for group quantization. """ + category = "GroupQuant" + bwd_category = None + def __init__(self, event, arch=None, python_path=None): self.event = event self.arch = arch diff --git a/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py b/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py index fb0284e2f..b6cac45a3 100644 --- a/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py @@ -30,6 +30,9 @@ class aiter_rms_norm(RMSNorm): flops/bytes are inherited from RMSNorm (affine=True, training=False). """ + category = "RMSNorm" + bwd_category = None + def __init__(self, event, arch=None, python_path=None): # Normalization.__init__ calls self.get_param_details and sets all attrs super().__init__(event, arch, python_path) @@ -79,6 +82,9 @@ class aiter_rmsnorm(RMSNorm): get_param_details uses input at index [1] and weight length at [2][0]. """ + category = "RMSNorm" + bwd_category = None + @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][1]) # input: [M, N] @@ -118,6 +124,9 @@ class aiter_rmsnorm2d_fwd_with_dynamicquant_ck(RMSNorm): Bytes: read input+weight, write out (FP8) + yscale (FP32). """ + category = "RMSNorm" + bwd_category = None + @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][1]) # input: [M, N] @@ -174,6 +183,9 @@ def __init__(self, event, arch=None, python_path=None): super().__init__(event, arch, python_path) self.group_size = int(event["args"]["Concrete Inputs"][3]) + category = "RMSNorm" + bwd_category = None + @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][0]) @@ -228,6 +240,9 @@ class aiter_rmsnorm2d_fwd_with_add_ck(RMSNorm): Bytes: HBM traffic per GPU (read input+residual_in+weight, write out+residual_out). """ + category = "RMSNorm" + bwd_category = None + @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][1]) # input: [M, N] @@ -311,6 +326,9 @@ def __init__(self, event, arch=None, python_path=None): super().__init__(event, arch, python_path) self.group_size = int(event["args"]["Concrete Inputs"][4]) + category = "RMSNorm" + bwd_category = None + @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][0]) @@ -404,6 +422,9 @@ def __init__(self, event, arch=None, python_path=None): else: self.n_out = N + category = "RMSNorm" + bwd_category = None + @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][0]) diff --git a/TraceLens/PerfModel/op_categories.py b/TraceLens/PerfModel/op_categories.py index f64372e0f..c6fe893d3 100644 --- a/TraceLens/PerfModel/op_categories.py +++ b/TraceLens/PerfModel/op_categories.py @@ -4,13 +4,12 @@ # See LICENSE for license information. ############################################################################### -""" -Registry-based categorization of CPU torch ops. +"""Registry-based categorization of CPU torch ops. -The source of truth is a flat ``op_name -> final category`` registry used by -``categorize_torch_op``. ``dict_cat2names`` remains as a compatibility view for -legacy per-category sheets such as ``GEMM``, ``CONV_fwd`` / ``CONV_bwd``, and -``GroupedGEMM_fwd`` / ``GroupedGEMM_bwd``. +Perf model classes declare their own output categories via ``category`` and, +when linked-forward backward metrics are intentionally supported, +``bwd_category``. Category-only ops that do not have a perf model live in +``CATEGORY_ONLY_OP_CATEGORIES``. """ import re @@ -19,27 +18,11 @@ from typing import Optional, Pattern, Tuple -SDPA_BWD_OPS = frozenset( - { - "FlashAttnFuncBackward", - "FusedAttnFuncBackward", - "flash_attn::_flash_attn_backward", - "flash_attn::_flash_attn_varlen_backward", - "aten::_scaled_dot_product_cudnn_attention_backward", - "aten::_scaled_dot_product_efficient_attention_backward", - "aten::_scaled_dot_product_flash_attention_backward", - "aiter::_flash_attn_backward", - "aiter::wrapper_fmha_v3_bwd", - "aiter::mha_bwd", - } -) - - -OP_CATEGORY_OVERRIDES: Dict[str, str] = { +CATEGORY_ONLY_OP_CATEGORIES: Dict[str, str] = { # CONV ops not present in op_to_perf_model_class_map. "aten::miopen_convolution": "CONV_fwd", "aten::cudnn_convolution": "CONV_fwd", - # SDPA backward ops that do not all have perf models in core TraceLens. + # SDPA backward ops without direct perf models in core TraceLens. "FlashAttnFuncBackward": "SDPA_bwd", "FusedAttnFuncBackward": "SDPA_bwd", "aten::_scaled_dot_product_cudnn_attention_backward": "SDPA_bwd", @@ -119,33 +102,20 @@ def is_backward_op(op_name: str) -> bool: op_name.endswith("_backward") or op_name.endswith("Backward") or op_name.endswith("_bwd") - or op_name in SDPA_BWD_OPS ) -def resolve_base_category(op_name: str, base_category: str) -> str: - """Return the final category for an op in a base/sheet category.""" - if base_category == "SDPA": - return "SDPA_bwd" if is_backward_op(op_name) else "SDPA_fwd" - if base_category == "Normalization": - return "NORM_bwd" if is_backward_op(op_name) else "NORM_fwd" - if base_category == "CONV": - return "CONV_bwd" if is_backward_op(op_name) else "CONV_fwd" - if base_category == "GroupedGEMM": - return "GroupedGEMM_bwd" if is_backward_op(op_name) else "GroupedGEMM_fwd" - if base_category == "SSM": - return "SSM_bwd" if is_backward_op(op_name) else "SSM_fwd" - if base_category == "MoE_comm": - return "MoE_comm_bwd" if is_backward_op(op_name) else "MoE_comm_fwd" - if base_category == "RoPE": - return "RoPE_bwd" if is_backward_op(op_name) else "RoPE_fwd" - if base_category == "CrossEntropy": - return "CrossEntropy_bwd" if is_backward_op(op_name) else "CrossEntropy_fwd" - if base_category in ("BinaryElementwise", "UnaryElementwise"): - return "elementwise" - if base_category == "Reduce": - return "reduce" - return base_category +def get_perf_model_category(perf_model_class: type, bwd: bool = False) -> Optional[str]: + """Return the category declared by a perf model class.""" + attr_name = "bwd_category" if bwd else "category" + category = getattr(perf_model_class, attr_name, None) + if bwd: + return category + if category is None: + raise ValueError( + f"perf_model_class {perf_model_class} must define a category attribute" + ) + return category def sheet_category_from_final_category(category: str) -> str: @@ -173,88 +143,93 @@ def sheet_category_from_final_category(category: str) -> str: return category_to_sheet.get(category, category) +def category_from_sheet_category(op_name: str, sheet_category: str) -> str: + """Resolve a legacy sheet category to the final categorization label.""" + sheet_to_fwd_bwd = { + "CONV": ("CONV_fwd", "CONV_bwd"), + "SDPA": ("SDPA_fwd", "SDPA_bwd"), + "Normalization": ("NORM_fwd", "NORM_bwd"), + "GroupedGEMM": ("GroupedGEMM_fwd", "GroupedGEMM_bwd"), + "SSM": ("SSM_fwd", "SSM_bwd"), + "MoE_comm": ("MoE_comm_fwd", "MoE_comm_bwd"), + "RoPE": ("RoPE_fwd", "RoPE_bwd"), + "CrossEntropy": ("CrossEntropy_fwd", "CrossEntropy_bwd"), + } + if sheet_category in sheet_to_fwd_bwd: + fwd_category, bwd_category = sheet_to_fwd_bwd[sheet_category] + return bwd_category if is_backward_op(op_name) else fwd_category + if sheet_category in ("BinaryElementwise", "UnaryElementwise"): + return "elementwise" + if sheet_category == "Reduce": + return "reduce" + return sheet_category + + def category_from_sheet_view( op_name: str, dict_cat2names: Mapping[str, List[str]] ) -> Optional[str]: """Compatibility fallback for callers that still mutate ``dict_cat2names``.""" for sheet_category, names in dict_cat2names.items(): if op_name in names: - return resolve_base_category(op_name, sheet_category) + return category_from_sheet_category(op_name, sheet_category) return None def build_op_category_registry( op_to_perf_model_class_map: Mapping[str, type], - dict_base_class2category: Mapping[type, str], - overrides: Optional[Mapping[str, str]] = None, + category_only_ops: Optional[Mapping[str, str]] = None, ) -> Dict[str, str]: """Construct the flat ``op_name -> final category`` registry.""" registry: Dict[str, str] = {} for op_name, perf_model_class in op_to_perf_model_class_map.items(): - base_classes = perf_model_class.__bases__ - if len(base_classes) != 1: - raise ValueError( - f"op_name: {op_name}, perf_model_class: {perf_model_class}, " - f"base_classes: {base_classes}" - ) - base_category = dict_base_class2category.get(base_classes[0]) - if base_category is None: - raise ValueError( - f"op_name: {op_name}, perf_model_class: {perf_model_class}, " - f"base_class: {base_classes[0]}" - ) - registry[op_name] = resolve_base_category(op_name, base_category) - - if overrides: - registry.update(overrides) + registry[op_name] = get_perf_model_category(perf_model_class) + + if category_only_ops: + registry.update(category_only_ops) return registry def build_dict_cat2names( op_to_perf_model_class_map: Mapping[str, type], - dict_base_class2category: Mapping[type, str], ) -> DefaultDict[str, List[str]]: """Build the legacy ``category -> op names`` view used for report sheets.""" dict_cat2names = defaultdict(list) # type: DefaultDict[str, List[str]] for op_name, perf_model_class in op_to_perf_model_class_map.items(): + category = get_perf_model_category(perf_model_class) + dict_cat2names[sheet_category_from_final_category(category)].append(op_name) + return dict_cat2names + + +def build_dict_base_class2category( + op_to_perf_model_class_map: Mapping[str, type], +) -> Dict[type, str]: + """Compatibility view for callers that still inspect base-class categories.""" + base_class2category: Dict[type, str] = {} + for perf_model_class in op_to_perf_model_class_map.values(): base_classes = perf_model_class.__bases__ if len(base_classes) != 1: - raise ValueError( - f"op_name: {op_name}, perf_model_class: {perf_model_class}, " - f"base_classes: {base_classes}" - ) - base_category = dict_base_class2category.get(base_classes[0]) - if base_category is None: - raise ValueError( - f"op_name: {op_name}, perf_model_class: {perf_model_class}, " - f"base_class: {base_classes[0]}" - ) - dict_cat2names[base_category].append(op_name) - return dict_cat2names + continue + base_class = base_classes[0] + category = get_perf_model_category(perf_model_class) + sheet_category = sheet_category_from_final_category(category) + existing = base_class2category.get(base_class) + if existing is not None and existing != sheet_category: + continue + base_class2category[base_class] = sheet_category + return base_class2category def register_perf_model_categories( perf_model_extension: Mapping[str, type], - dict_base_class2category: Mapping[type, str], registry: MutableMapping[str, str], dict_cat2names: MutableMapping[str, List[str]], ) -> None: """Register categories for extension-provided perf models.""" for op_name, perf_model_class in perf_model_extension.items(): - base_classes = perf_model_class.__bases__ - if len(base_classes) != 1: - raise ValueError( - f"op_name: {op_name}, perf_model_class: {perf_model_class}, " - f"base_classes: {base_classes}" - ) - sheet_category = dict_base_class2category.get(base_classes[0]) - if sheet_category is None: - raise ValueError( - f"op_name: {op_name}, perf_model_class: {perf_model_class}, " - f"base_class: {base_classes[0]}" - ) - registry[op_name] = resolve_base_category(op_name, sheet_category) + category = get_perf_model_category(perf_model_class) + sheet_category = sheet_category_from_final_category(category) + registry[op_name] = category if sheet_category not in dict_cat2names: dict_cat2names[sheet_category] = [] _append_unique(dict_cat2names[sheet_category], [op_name]) @@ -263,33 +238,9 @@ def register_perf_model_categories( def register_op_categories( op_category_extension: Mapping[str, str], registry: MutableMapping[str, str], - dict_cat2names: Optional[MutableMapping[str, List[str]]] = None, ) -> None: """Register explicit category-only op labels.""" registry.update(op_category_extension) - if dict_cat2names is None: - return - for op_name, category in op_category_extension.items(): - sheet_category = sheet_category_from_final_category(category) - if sheet_category not in dict_cat2names: - dict_cat2names[sheet_category] = [] - _append_unique(dict_cat2names[sheet_category], [op_name]) - - -def register_dict_cat2names_extension( - dict_cat2names_extension: Mapping[str, List[str]], - registry: MutableMapping[str, str], - dict_cat2names: MutableMapping[str, List[str]], -) -> None: - """Support the older ``dict_cat2names_extension`` extension contract.""" - for sheet_category, names in dict_cat2names_extension.items(): - if not isinstance(names, list): - raise ValueError(f"Expected names to be a list, got {type(names)}") - if sheet_category not in dict_cat2names: - dict_cat2names[sheet_category] = [] - _append_unique(dict_cat2names[sheet_category], names) - for op_name in names: - registry[op_name] = resolve_base_category(op_name, sheet_category) def categorize_torch_op_from_registry( diff --git a/TraceLens/PerfModel/perf_model.py b/TraceLens/PerfModel/perf_model.py index a0060aea4..64467c034 100644 --- a/TraceLens/PerfModel/perf_model.py +++ b/TraceLens/PerfModel/perf_model.py @@ -25,6 +25,8 @@ class GEMM: If you want to add a new GEMM operation, you should inherit from this class. """ + category = "GEMM" + bwd_category = None cache_gemm_results = {} # This is used to cache gemm results _origami_import_error_printed = False @@ -862,6 +864,9 @@ def bytes_bwd(self, bytes_per_element): class CONV: # Conv perf model is based on: https://github.com/pytorch/pytorch/blob/main/torch/utils/flop_counter.py # we will make stuff reusiable across conv1d, conv2d, and conv3d + category = "CONV_fwd" + bwd_category = "CONV_bwd" + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.param_details = self.get_param_details(event) @@ -1150,6 +1155,8 @@ def bytes_bwd(self): class aten_conv_bwd(CONV): + category = "CONV_bwd" + @staticmethod def get_param_details(event): # convolution_backward signature: @@ -1374,6 +1381,8 @@ class ConvBias_Backward(CONV): Uses cached forward pass parameters via sequence number linkage. """ + category = "CONV_bwd" + @staticmethod def get_param_details(event): # Try to get forward pass parameters using sequence number @@ -1583,6 +1592,8 @@ class ConvBiasReLU_Backward(CONV): ReLU backward: gradient is masked where forward output was negative. """ + category = "CONV_bwd" + @staticmethod def get_param_details(event): # Try to get forward pass parameters using sequence number @@ -1746,6 +1757,8 @@ def get_time( # 4. Scaled Dot Product Attention class SDPA: + category = "SDPA_fwd" + bwd_category = "SDPA_bwd" def __init__(self, event, arch=None, python_path=None, enable_origami=False): # S = QK^T @@ -2308,6 +2321,8 @@ def get_param_details(event): class flash_attention_backward(SDPA): """Backward pass for flash_attn::_flash_attn_backward. Argument order: dout, q, k, v, ...""" + category = "SDPA_bwd" + def __init__(self, event, arch=None, python_path=None, enable_origami=False): super().__init__(event, arch, python_path, enable_origami=enable_origami) self.d_h = ( @@ -2471,6 +2486,8 @@ def flops(self): class flash_attention_varlen_backward(SDPA): + category = "SDPA_bwd" + def __init__(self, event, arch=None, python_path=None, enable_origami=False): super().__init__(event, arch, python_path, enable_origami=enable_origami) self.num_seqs_q, self.num_seqs_kv, self.max_seqlen_q, self.max_seqlen_kv = ( @@ -2792,6 +2809,7 @@ def get_param_details(event): class aiter__flash_attn_backward(SDPA): + category = "SDPA_bwd" @staticmethod def get_param_details(event): @@ -2923,6 +2941,7 @@ def get_param_details(event): class aiter__fmha_v3_backward(SDPA): + category = "SDPA_bwd" @staticmethod def get_param_details(event): @@ -3030,6 +3049,8 @@ class aiter__mha_bwd(SDPA): # aiter::mha_bwd(dout, q, k, v, out, softmax_lse, dropout_p, softmax_scale, is_causal, ...) # q[1], k[2], v[3] — raw shape (B, N, H, d_h) in bnhd order; bhnd_idx=(0,2,1,3) extracts B,H,N,d_h. + category = "SDPA_bwd" + @staticmethod def get_param_details(event): input_dims = event["args"]["Input Dims"] @@ -3202,6 +3223,8 @@ def get_param_details(event): class UnaryElementwise: + category = "elementwise" + bwd_category = None def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event @@ -3268,6 +3291,8 @@ def get_param_details(event): class BinaryElementwise: + category = "elementwise" + bwd_category = None def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event @@ -3409,6 +3434,9 @@ class Reduce: Models reduction over one or more dimensions of a tensor. """ + category = "reduce" + bwd_category = None + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.arch = arch @@ -3655,6 +3683,9 @@ class GroupedGemm: - Total bytes: (2*M*N) * bpe_out + (2*M*K) * bpe_in + (2*G*K*N) * bpe_in """ + category = "GroupedGEMM_fwd" + bwd_category = "GroupedGEMM_bwd" + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.param_details = self.get_param_details(event) @@ -4226,6 +4257,9 @@ def parse_list(input: str, dtype): class Normalization: + category = "NORM_fwd" + bwd_category = "NORM_bwd" + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.arch = arch @@ -4497,6 +4531,8 @@ def bytes_bwd(self, bytes_per_element): class BatchNormBwd(Normalization): + category = "NORM_bwd" + @staticmethod def get_param_details(event): # miopen_batch_norm_backward and cudnn_batch_norm_backward have different paramters @@ -4605,6 +4641,8 @@ def bytes_bwd(self, bytes_per_element): class LayerNormBwd(Normalization): + category = "NORM_bwd" + @staticmethod def get_param_details(event): args_input_dims = event["args"]["Input Dims"] @@ -4703,6 +4741,8 @@ def bytes_bwd(self, bytes_per_element): class GroupNormBwd(Normalization): + category = "NORM_bwd" + @staticmethod def get_param_details(event): args_input_dims = event["args"]["Input Dims"] @@ -4796,6 +4836,8 @@ def bytes_bwd(self, bytes_per_element): class InstanceNormBwd(Normalization): # instance_norm_backward does not appear in the traces and does not even exist in # https://github.com/pytorch/pytorch/blob/1457786f7445fb0e72794aa98c0ebaa3bc24ced5/aten/src/ATen/native/Normalization.cpp + category = "NORM_bwd" + @staticmethod def get_param_details(event): raise NotImplementedError(f"Backward pass for InstanceNorm is not defined.") @@ -4866,6 +4908,8 @@ def bytes_bwd(self, bytes_per_element): class RMSNormBwd(Normalization): + category = "NORM_bwd" + @staticmethod def get_param_details(event): args_input_dims = event["args"]["Input Dims"] @@ -4945,6 +4989,9 @@ class MoEComm: bytes() = num_tokens × hidden_dim × bpe (data volume moved). """ + category = "MoE_comm_fwd" + bwd_category = "MoE_comm_bwd" + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.arch = arch @@ -5015,6 +5062,9 @@ class CausalConv1d: FLOPS = 2 × batch × channels × seq_len × kernel_size (depthwise conv). """ + category = "SSM_fwd" + bwd_category = None + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.arch = arch @@ -5096,6 +5146,9 @@ class FusedRoPE: Total = 3 × num_elements (since 6 ops per 2 elements). """ + category = "RoPE_fwd" + bwd_category = None + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.arch = arch @@ -5158,6 +5211,9 @@ class CrossEntropy: FLOPS ≈ 5 × batch × vocab_size (exp + sum + log + subtract + lookup per element). """ + category = "CrossEntropy_fwd" + bwd_category = None + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.arch = arch @@ -5246,6 +5302,9 @@ class MambaSSD: C@h (step 4): 2 · B · H · T · N · P """ + category = "SSM_fwd" + bwd_category = "SSM_bwd" + def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event self.arch = arch diff --git a/TraceLens/PerfModel/torch_op_mapping.py b/TraceLens/PerfModel/torch_op_mapping.py index d063af6d1..134897446 100644 --- a/TraceLens/PerfModel/torch_op_mapping.py +++ b/TraceLens/PerfModel/torch_op_mapping.py @@ -5,9 +5,10 @@ ############################################################################### from . import perf_model -from .extensions import get_pseudo_op_mappings, get_pseudo_op_categories +from .extensions import get_pseudo_op_mappings from .op_categories import ( - OP_CATEGORY_OVERRIDES, + CATEGORY_ONLY_OP_CATEGORIES, + build_dict_base_class2category, build_dict_cat2names, build_op_category_registry, categorize_torch_op_from_registry, @@ -200,34 +201,15 @@ for op in reduce_ops: op_to_perf_model_class_map[op] = perf_model.aten_reduce -dict_base_class2category = { - perf_model.GEMM: "GEMM", - perf_model.GroupedGemm: "GroupedGEMM", - perf_model.CONV: "CONV", - perf_model.SDPA: "SDPA", - perf_model.UnaryElementwise: "UnaryElementwise", - perf_model.BinaryElementwise: "BinaryElementwise", - perf_model.Normalization: "Normalization", - perf_model.Reduce: "Reduce", - perf_model.MoEComm: "MoE_comm", - perf_model.CausalConv1d: "SSM", - perf_model.FusedRoPE: "RoPE", - perf_model.CrossEntropy: "CrossEntropy", - perf_model.MambaSSD: "SSM", -} - -# Add pseudo-op extension categories -dict_base_class2category.update(get_pseudo_op_categories()) +# Compatibility view for older callers that inspect base-class categories. New +# categorization reads ``category`` / ``bwd_category`` from perf model classes. +dict_base_class2category = build_dict_base_class2category(op_to_perf_model_class_map) -dict_cat2names = build_dict_cat2names( - op_to_perf_model_class_map, - dict_base_class2category, -) +dict_cat2names = build_dict_cat2names(op_to_perf_model_class_map) OP_CATEGORY_REGISTRY = build_op_category_registry( op_to_perf_model_class_map, - dict_base_class2category, - overrides=OP_CATEGORY_OVERRIDES, + category_only_ops=CATEGORY_ONLY_OP_CATEGORIES, ) diff --git a/TraceLens/Reporting/generate_perf_report_pytorch.py b/TraceLens/Reporting/generate_perf_report_pytorch.py index fbc46327a..bdb23dbcd 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch.py @@ -120,14 +120,10 @@ def apply_extension(perf_analyzer, extension_path): extension_name = os.path.splitext(os.path.basename(extension_path))[0] from TraceLens.PerfModel.op_categories import ( - register_dict_cat2names_extension, register_op_categories, register_perf_model_categories, ) - from TraceLens.PerfModel.torch_op_mapping import ( - OP_CATEGORY_REGISTRY, - dict_base_class2category, - ) + from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY spec = importlib.util.spec_from_file_location(extension_name, extension_path) extension = importlib.util.module_from_spec(spec) @@ -149,7 +145,6 @@ def apply_extension(perf_analyzer, extension_path): perf_analyzer.op_to_perf_model_class_map.update(perf_model_extension) register_perf_model_categories( perf_model_extension, - dict_base_class2category, OP_CATEGORY_REGISTRY, perf_analyzer.dict_cat2names, ) @@ -163,19 +158,12 @@ def apply_extension(perf_analyzer, extension_path): register_op_categories( op_category_extension, OP_CATEGORY_REGISTRY, - perf_analyzer.dict_cat2names, ) if hasattr(extension, "dict_cat2names_extension"): - print(f"Updating dict_cat2names with extension from {extension_path}") - if not isinstance(extension.dict_cat2names_extension, dict): - raise ValueError( - f"Expected dict_cat2names_extension to be a dict, got {type(extension.dict_cat2names_extension)}" - ) - - register_dict_cat2names_extension( - extension.dict_cat2names_extension, - OP_CATEGORY_REGISTRY, - perf_analyzer.dict_cat2names, + warnings.warn( + "dict_cat2names_extension is deprecated and ignored. Use " + "perf_model_extension for modeled ops or op_category_extension for " + "category-only ops." ) diff --git a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py index 8dc125e1d..a31e9acb6 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py @@ -401,14 +401,10 @@ def apply_extension(perf_analyzer, extension_path): extension_name = os.path.splitext(os.path.basename(extension_path))[0] from TraceLens.PerfModel.op_categories import ( - register_dict_cat2names_extension, register_op_categories, register_perf_model_categories, ) - from TraceLens.PerfModel.torch_op_mapping import ( - OP_CATEGORY_REGISTRY, - dict_base_class2category, - ) + from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY spec = importlib.util.spec_from_file_location(extension_name, extension_path) extension = importlib.util.module_from_spec(spec) @@ -430,7 +426,6 @@ def apply_extension(perf_analyzer, extension_path): perf_analyzer.op_to_perf_model_class_map.update(perf_model_extension) register_perf_model_categories( perf_model_extension, - dict_base_class2category, OP_CATEGORY_REGISTRY, perf_analyzer.dict_cat2names, ) @@ -444,19 +439,12 @@ def apply_extension(perf_analyzer, extension_path): register_op_categories( op_category_extension, OP_CATEGORY_REGISTRY, - perf_analyzer.dict_cat2names, ) if hasattr(extension, "dict_cat2names_extension"): - print(f"Updating dict_cat2names with extension from {extension_path}") - if not isinstance(extension.dict_cat2names_extension, dict): - raise ValueError( - f"Expected dict_cat2names_extension to be a dict, got {type(extension.dict_cat2names_extension)}" - ) - - register_dict_cat2names_extension( - extension.dict_cat2names_extension, - OP_CATEGORY_REGISTRY, - perf_analyzer.dict_cat2names, + warnings.warn( + "dict_cat2names_extension is deprecated and ignored. Use " + "perf_model_extension for modeled ops or op_category_extension for " + "category-only ops." ) diff --git a/TraceLens/TreePerf/tree_perf.py b/TraceLens/TreePerf/tree_perf.py index ad0c7be33..382237e6a 100644 --- a/TraceLens/TreePerf/tree_perf.py +++ b/TraceLens/TreePerf/tree_perf.py @@ -29,6 +29,7 @@ dict_cat2names, op_to_perf_model_class_map, ) +from ..PerfModel.op_categories import get_perf_model_category from ..Trace2Tree.extensions import apply_pseudo_op_extensions from ..Trace2Tree.trace_capture_merge_experimental import merge_capture_trace_into_graph from ..Trace2Tree.trace_to_tree import JaxTraceToTree, TraceToTree @@ -1664,6 +1665,10 @@ def _is_sole_bwd_with_fwd_perf_model(self, event): if not fwd_event or not is_sole_bwd: return False + fwd_perf_model_class = self.op_to_perf_model_class_map.get(fwd_event["name"]) + if get_perf_model_category(fwd_perf_model_class, bwd=True) is None: + return False + # Check if backward metrics are actually defined for this forward op try: self.compute_perf_metrics(fwd_event, bwd=True) @@ -1925,6 +1930,21 @@ def list_to_tuple(obj): args = event.get("args", {}) has_own_perf_model = self._has_perf_model(event) is_sole_bwd = self._is_sole_bwd_with_fwd_perf_model(event) + linked_fwd_event = None + linked_bwd_category = None + if is_sole_bwd: + linked_fwd_event, _ = self._get_linked_fwd_event(event) + if linked_fwd_event is not None: + linked_perf_model_class = self.op_to_perf_model_class_map.get( + linked_fwd_event["name"] + ) + linked_bwd_category = get_perf_model_category( + linked_perf_model_class, bwd=True + ) + + op_category = self.op_categorizer(event) + if not has_own_perf_model and linked_bwd_category is not None: + op_category = linked_bwd_category if event.get("overlap_pct") is None and event.get("gpu_events"): kernels = [ @@ -1939,7 +1959,7 @@ def list_to_tuple(obj): row = { "name": event.get("name"), - "op category": self.op_categorizer(event), + "op category": op_category, "UID": event.get("UID"), "pid": event.get("pid"), "tid": event.get("tid"), @@ -1955,6 +1975,13 @@ def list_to_tuple(obj): "External id": args.get("External id"), "duration_us": event.get("dur"), "has_perf_model": has_own_perf_model or is_sole_bwd, + "metrics_source": ( + "own_perf_model" + if has_own_perf_model + else "linked_forward_bwd" + if is_sole_bwd + else "kernel_time_only" + ), "overlapping_kernel_names": event.get("overlapping_kernel_names"), "overlapping_kernels_details": event.get("overlapping_kernels_details"), "overlap_pct": event.get("overlap_pct"), @@ -2041,9 +2068,8 @@ def list_to_tuple(obj): ).compute_metrics()["busy_time"] elif include_perf_metrics and is_sole_bwd: # 1:1 backward op - use forward's backward metrics - fwd_event, _ = self._get_linked_fwd_event(event) try: - metrics = self.compute_perf_metrics(fwd_event, bwd=True) + metrics = self.compute_perf_metrics(linked_fwd_event, bwd=True) for col in perf_cols: if col in metrics: row[col] = metrics[col] @@ -2120,6 +2146,7 @@ def list_to_tuple(obj): ["Input Dims", "Input type", "Input Strides", "Concrete Inputs"] ) col_order.extend(["duration_us", "has_perf_model", "is_recompute"]) + col_order.append("metrics_source") if include_perf_metrics: col_order.extend(perf_cols) col_order.append("perf_params") @@ -2168,6 +2195,7 @@ def summarize_df_unified_perf_table( grouping_cols = [ "name", "op category", + "metrics_source", "process_name", "process_label", "thread_name", diff --git a/docs/generate_perf_report.md b/docs/generate_perf_report.md index 9c605292f..d5ee1b97e 100644 --- a/docs/generate_perf_report.md +++ b/docs/generate_perf_report.md @@ -164,7 +164,7 @@ Pass a Python file path via `--extension_file`. The file can define one or more |-----------------------------|-----------|-----------------------------------------------------------------------------| | `tree_postprocess_extension`| `Callable`| Called with `perf_analyzer.tree`. Use to modify the tree structure post-construction. | | `perf_model_extension` | `dict` | A mapping from op name (str) to a custom performance model class. These will override or extend existing models. | -| `dict_cat2names_extension` | `dict` | Mapping from new category names to lists of op names, merged into the built-in op categories. | +| `op_category_extension` | `dict` | Mapping from category-only op names to final categories, used when an op should appear in unified reports without a perf model. | #### 📄 Example Extension File for MegatronLM in the examples dir diff --git a/examples/example_megatron_extension.py b/examples/example_megatron_extension.py index edee93762..da1d16035 100644 --- a/examples/example_megatron_extension.py +++ b/examples/example_megatron_extension.py @@ -542,6 +542,8 @@ class te_layer_norm_fwd(Normalization): args[2] = ln_bias (beta, may be empty) """ + category = "NORM_fwd" + @staticmethod def get_param_details(event): args_input_dims = event["args"]["Input Dims"] @@ -579,6 +581,8 @@ class te_layer_norm_bwd(Normalization): which infers it from the (often empty) ln_bias dim. """ + category = "NORM_bwd" + @staticmethod def get_param_details(event): args_input_dims = event["args"]["Input Dims"] @@ -670,16 +674,7 @@ def get_param_details(event): "LayerNormFnBackward": te_layer_norm_bwd, } -dict_cat2names_extension = { - "GEMM": [ - "_Linear_yfwd_mm", - "_LinearBackward_xgrad_mm", - "_LinearBackward_wgrad_mm", - "_LayerNormLinear_yfwd_mm", - "_LayerNormLinearBackward_xgrad_mm", - "_LayerNormLinearBackward_wgrad_mm", - ], - "SDPA": ["FusedAttnFunc", "FusedAttnFuncBackward"], - "GroupedGEMM": ["GroupedGemm", "GroupedGemmBackward"], - "Normalization": ["LayerNormFn", "LayerNormFnBackward"], +op_category_extension = { + "FusedAttnFuncBackward": "SDPA_bwd", + "GroupedGemmBackward": "GroupedGEMM_bwd", } diff --git a/tests/test_pseudo_ops_extension.py b/tests/test_pseudo_ops_extension.py index f8eb3c514..5062cf206 100644 --- a/tests/test_pseudo_ops_extension.py +++ b/tests/test_pseudo_ops_extension.py @@ -28,7 +28,7 @@ tree_postprocess_extension, _link_checkpoint_fwd_bwd, perf_model_extension, - dict_cat2names_extension, + op_category_extension, te_layer_norm_fwd, te_layer_norm_bwd, ) @@ -540,41 +540,34 @@ def test_pseudo_ops_in_ops_summary(self): class TestFusedAttnFuncBackwardCategorization: """Test that FusedAttnFuncBackward is categorized as SDPA_bwd.""" - def test_categorization_via_dict_cat2names(self): - """FusedAttnFuncBackward must be in SDPA category.""" - assert "FusedAttnFuncBackward" in dict_cat2names_extension["SDPA"] + def test_categorization_via_op_category_extension(self): + """FusedAttnFuncBackward must be registered as a category-only op.""" + assert op_category_extension["FusedAttnFuncBackward"] == "SDPA_bwd" def test_categorize_as_sdpa_bwd(self): """Core categorizer must return SDPA_bwd for FusedAttnFuncBackward.""" - from TraceLens.PerfModel.torch_op_mapping import ( - categorize_torch_op, - dict_cat2names, - ) + from TraceLens.PerfModel.torch_op_mapping import categorize_torch_op - dict_cat2names["SDPA"].extend(dict_cat2names_extension["SDPA"]) - try: - result = categorize_torch_op({"name": "FusedAttnFuncBackward"}) - assert result == "SDPA_bwd", f"Expected SDPA_bwd, got {result}" - finally: - for name in dict_cat2names_extension["SDPA"]: - if name in dict_cat2names["SDPA"]: - dict_cat2names["SDPA"].remove(name) + result = categorize_torch_op({"name": "FusedAttnFuncBackward"}) + assert result == "SDPA_bwd", f"Expected SDPA_bwd, got {result}" def test_fused_attn_fwd_still_sdpa_fwd(self): """FusedAttnFunc (forward) must still be SDPA_fwd.""" + from TraceLens.PerfModel.op_categories import register_perf_model_categories from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, categorize_torch_op, dict_cat2names, ) - dict_cat2names["SDPA"].extend(dict_cat2names_extension["SDPA"]) - try: - result = categorize_torch_op({"name": "FusedAttnFunc"}) - assert result == "SDPA_fwd", f"Expected SDPA_fwd, got {result}" - finally: - for name in dict_cat2names_extension["SDPA"]: - if name in dict_cat2names["SDPA"]: - dict_cat2names["SDPA"].remove(name) + register_perf_model_categories( + {"FusedAttnFunc": perf_model_extension["FusedAttnFunc"]}, + OP_CATEGORY_REGISTRY, + dict_cat2names, + ) + + result = categorize_torch_op({"name": "FusedAttnFunc"}) + assert result == "SDPA_fwd", f"Expected SDPA_fwd, got {result}" class TestLayerNormFnPerfModel: @@ -648,27 +641,30 @@ def test_layer_norm_fn_bwd_consistent_with_fwd(self): ), f"is_affine mismatch: fwd={fwd_model.is_affine}, bwd={bwd_model.is_affine}" def test_categorization_normalization(self): - """LayerNormFn/LayerNormFnBackward must be in Normalization category.""" - assert "LayerNormFn" in dict_cat2names_extension["Normalization"] - assert "LayerNormFnBackward" in dict_cat2names_extension["Normalization"] + """LayerNormFn classes declare their categories directly.""" + assert te_layer_norm_fwd.category == "NORM_fwd" + assert te_layer_norm_bwd.category == "NORM_bwd" def test_categorize_as_norm_fwd_bwd(self): """Core categorizer must return NORM_fwd and NORM_bwd.""" + from TraceLens.PerfModel.op_categories import register_perf_model_categories from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, categorize_torch_op, dict_cat2names, ) - dict_cat2names["Normalization"].extend( - dict_cat2names_extension["Normalization"] + register_perf_model_categories( + { + "LayerNormFn": te_layer_norm_fwd, + "LayerNormFnBackward": te_layer_norm_bwd, + }, + OP_CATEGORY_REGISTRY, + dict_cat2names, ) - try: - assert categorize_torch_op({"name": "LayerNormFn"}) == "NORM_fwd" - assert categorize_torch_op({"name": "LayerNormFnBackward"}) == "NORM_bwd" - finally: - for name in dict_cat2names_extension["Normalization"]: - if name in dict_cat2names["Normalization"]: - dict_cat2names["Normalization"].remove(name) + + assert categorize_torch_op({"name": "LayerNormFn"}) == "NORM_fwd" + assert categorize_torch_op({"name": "LayerNormFnBackward"}) == "NORM_bwd" def test_perf_model_extension_registration(self): """LayerNormFn/LayerNormFnBackward must be registered in perf_model_extension.""" diff --git a/tests/test_torch_op_categorization_registry.py b/tests/test_torch_op_categorization_registry.py index 24ec031ac..861973287 100644 --- a/tests/test_torch_op_categorization_registry.py +++ b/tests/test_torch_op_categorization_registry.py @@ -12,7 +12,6 @@ from TraceLens.PerfModel.op_categories import ( category_from_sheet_view, - register_dict_cat2names_extension, register_op_categories, ) from TraceLens.PerfModel.torch_op_mapping import ( @@ -99,6 +98,7 @@ def test_dict_cat2names_keeps_sheet_compatibility_view(): assert "aten::convolution" in dict_cat2names["CONV"] assert "aten::convolution_backward" in dict_cat2names["CONV"] assert "FlashAttnFuncBackward" not in dict_cat2names["SDPA"] + assert "TokenPermuteMaskMap" not in dict_cat2names["MoE_comm"] def test_dict_cat2names_dynamic_fallback_for_legacy_callers(): @@ -109,30 +109,12 @@ def test_dict_cat2names_dynamic_fallback_for_legacy_callers(): ) -def test_register_op_category_extension_updates_registry_and_sheet_view(): +def test_register_op_category_extension_updates_registry_only(): registry = {} - sheet_view = {} register_op_categories( {"MyCategoryOnlyBackward": "SDPA_bwd"}, registry, - sheet_view, ) assert registry["MyCategoryOnlyBackward"] == "SDPA_bwd" - assert sheet_view["SDPA"] == ["MyCategoryOnlyBackward"] - - -def test_register_dict_cat2names_extension_updates_registry_and_sheet_view(): - registry = {} - sheet_view = {} - - register_dict_cat2names_extension( - {"GroupedGEMM": ["GroupedGemm", "GroupedGemmBackward"]}, - registry, - sheet_view, - ) - - assert registry["GroupedGemm"] == "GroupedGEMM_fwd" - assert registry["GroupedGemmBackward"] == "GroupedGEMM_bwd" - assert sheet_view["GroupedGEMM"] == ["GroupedGemm", "GroupedGemmBackward"] From bd8df052b7e93cae865c57b71c8d6c5582c62d55 Mon Sep 17 00:00:00 2001 From: Jassani Date: Fri, 8 May 2026 12:56:31 -0400 Subject: [PATCH 04/10] PerfModel: simplify category registry flow Keep categorization registry-only, retain legacy sheet categories only for report sheet generation, and move RMSNorm extension category metadata onto an extension-side base class. Co-authored-by: Cursor --- .../rmsnorm_perf_model_extensions.py | 30 +++-------- TraceLens/PerfModel/op_categories.py | 50 ++++--------------- TraceLens/PerfModel/perf_model.py | 2 + TraceLens/PerfModel/torch_op_mapping.py | 1 - .../test_torch_op_categorization_registry.py | 13 +---- 5 files changed, 20 insertions(+), 76 deletions(-) diff --git a/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py b/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py index b6cac45a3..89e2ec4d0 100644 --- a/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py @@ -8,7 +8,14 @@ Performance models for RMSNorm pseudo-op extensions. """ -from TraceLens.PerfModel.perf_model import RMSNorm +from TraceLens.PerfModel.perf_model import RMSNorm as CoreRMSNorm + + +class RMSNorm(CoreRMSNorm): + """Extension-side RMSNorm family reported separately from generic normalization.""" + + category = "RMSNorm" + bwd_category = None class aiter_rms_norm(RMSNorm): @@ -30,9 +37,6 @@ class aiter_rms_norm(RMSNorm): flops/bytes are inherited from RMSNorm (affine=True, training=False). """ - category = "RMSNorm" - bwd_category = None - def __init__(self, event, arch=None, python_path=None): # Normalization.__init__ calls self.get_param_details and sets all attrs super().__init__(event, arch, python_path) @@ -82,9 +86,6 @@ class aiter_rmsnorm(RMSNorm): get_param_details uses input at index [1] and weight length at [2][0]. """ - category = "RMSNorm" - bwd_category = None - @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][1]) # input: [M, N] @@ -124,9 +125,6 @@ class aiter_rmsnorm2d_fwd_with_dynamicquant_ck(RMSNorm): Bytes: read input+weight, write out (FP8) + yscale (FP32). """ - category = "RMSNorm" - bwd_category = None - @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][1]) # input: [M, N] @@ -183,9 +181,6 @@ def __init__(self, event, arch=None, python_path=None): super().__init__(event, arch, python_path) self.group_size = int(event["args"]["Concrete Inputs"][3]) - category = "RMSNorm" - bwd_category = None - @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][0]) @@ -240,9 +235,6 @@ class aiter_rmsnorm2d_fwd_with_add_ck(RMSNorm): Bytes: HBM traffic per GPU (read input+residual_in+weight, write out+residual_out). """ - category = "RMSNorm" - bwd_category = None - @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][1]) # input: [M, N] @@ -326,9 +318,6 @@ def __init__(self, event, arch=None, python_path=None): super().__init__(event, arch, python_path) self.group_size = int(event["args"]["Concrete Inputs"][4]) - category = "RMSNorm" - bwd_category = None - @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][0]) @@ -422,9 +411,6 @@ def __init__(self, event, arch=None, python_path=None): else: self.n_out = N - category = "RMSNorm" - bwd_category = None - @staticmethod def get_param_details(event): op_shape = tuple(event["args"]["Input Dims"][0]) diff --git a/TraceLens/PerfModel/op_categories.py b/TraceLens/PerfModel/op_categories.py index c6fe893d3..73571fe9d 100644 --- a/TraceLens/PerfModel/op_categories.py +++ b/TraceLens/PerfModel/op_categories.py @@ -143,36 +143,12 @@ def sheet_category_from_final_category(category: str) -> str: return category_to_sheet.get(category, category) -def category_from_sheet_category(op_name: str, sheet_category: str) -> str: - """Resolve a legacy sheet category to the final categorization label.""" - sheet_to_fwd_bwd = { - "CONV": ("CONV_fwd", "CONV_bwd"), - "SDPA": ("SDPA_fwd", "SDPA_bwd"), - "Normalization": ("NORM_fwd", "NORM_bwd"), - "GroupedGEMM": ("GroupedGEMM_fwd", "GroupedGEMM_bwd"), - "SSM": ("SSM_fwd", "SSM_bwd"), - "MoE_comm": ("MoE_comm_fwd", "MoE_comm_bwd"), - "RoPE": ("RoPE_fwd", "RoPE_bwd"), - "CrossEntropy": ("CrossEntropy_fwd", "CrossEntropy_bwd"), - } - if sheet_category in sheet_to_fwd_bwd: - fwd_category, bwd_category = sheet_to_fwd_bwd[sheet_category] - return bwd_category if is_backward_op(op_name) else fwd_category - if sheet_category in ("BinaryElementwise", "UnaryElementwise"): - return "elementwise" - if sheet_category == "Reduce": - return "reduce" - return sheet_category - - -def category_from_sheet_view( - op_name: str, dict_cat2names: Mapping[str, List[str]] -) -> Optional[str]: - """Compatibility fallback for callers that still mutate ``dict_cat2names``.""" - for sheet_category, names in dict_cat2names.items(): - if op_name in names: - return category_from_sheet_category(op_name, sheet_category) - return None +def get_perf_model_sheet_category(perf_model_class: type) -> str: + """Return the legacy sheet category for a perf model class.""" + sheet_category = getattr(perf_model_class, "sheet_category", None) + if sheet_category is not None: + return sheet_category + return sheet_category_from_final_category(get_perf_model_category(perf_model_class)) def build_op_category_registry( @@ -196,8 +172,7 @@ def build_dict_cat2names( """Build the legacy ``category -> op names`` view used for report sheets.""" dict_cat2names = defaultdict(list) # type: DefaultDict[str, List[str]] for op_name, perf_model_class in op_to_perf_model_class_map.items(): - category = get_perf_model_category(perf_model_class) - dict_cat2names[sheet_category_from_final_category(category)].append(op_name) + dict_cat2names[get_perf_model_sheet_category(perf_model_class)].append(op_name) return dict_cat2names @@ -211,8 +186,7 @@ def build_dict_base_class2category( if len(base_classes) != 1: continue base_class = base_classes[0] - category = get_perf_model_category(perf_model_class) - sheet_category = sheet_category_from_final_category(category) + sheet_category = get_perf_model_sheet_category(perf_model_class) existing = base_class2category.get(base_class) if existing is not None and existing != sheet_category: continue @@ -228,7 +202,7 @@ def register_perf_model_categories( """Register categories for extension-provided perf models.""" for op_name, perf_model_class in perf_model_extension.items(): category = get_perf_model_category(perf_model_class) - sheet_category = sheet_category_from_final_category(category) + sheet_category = get_perf_model_sheet_category(perf_model_class) registry[op_name] = category if sheet_category not in dict_cat2names: dict_cat2names[sheet_category] = [] @@ -246,7 +220,6 @@ def register_op_categories( def categorize_torch_op_from_registry( row, registry: Mapping[str, str], - dict_cat2names: Optional[Mapping[str, List[str]]] = None, patterns: Iterable[Tuple[Pattern, str]] = OP_CATEGORY_PATTERNS, ) -> str: """Return the category for ``row`` using explicit registry data.""" @@ -256,11 +229,6 @@ def categorize_torch_op_from_registry( if category is not None: return category - if dict_cat2names is not None: - category = category_from_sheet_view(name, dict_cat2names) - if category is not None: - return category - for pattern, category in patterns: if pattern.match(name): return category diff --git a/TraceLens/PerfModel/perf_model.py b/TraceLens/PerfModel/perf_model.py index 64467c034..a22949f89 100644 --- a/TraceLens/PerfModel/perf_model.py +++ b/TraceLens/PerfModel/perf_model.py @@ -3225,6 +3225,7 @@ def get_param_details(event): class UnaryElementwise: category = "elementwise" bwd_category = None + sheet_category = "UnaryElementwise" def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event @@ -3293,6 +3294,7 @@ def get_param_details(event): class BinaryElementwise: category = "elementwise" bwd_category = None + sheet_category = "BinaryElementwise" def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event diff --git a/TraceLens/PerfModel/torch_op_mapping.py b/TraceLens/PerfModel/torch_op_mapping.py index 134897446..9b1469c1a 100644 --- a/TraceLens/PerfModel/torch_op_mapping.py +++ b/TraceLens/PerfModel/torch_op_mapping.py @@ -237,5 +237,4 @@ def categorize_torch_op(row): return categorize_torch_op_from_registry( row, OP_CATEGORY_REGISTRY, - dict_cat2names=dict_cat2names, ) diff --git a/tests/test_torch_op_categorization_registry.py b/tests/test_torch_op_categorization_registry.py index 861973287..fa11838d6 100644 --- a/tests/test_torch_op_categorization_registry.py +++ b/tests/test_torch_op_categorization_registry.py @@ -10,10 +10,7 @@ import pytest -from TraceLens.PerfModel.op_categories import ( - category_from_sheet_view, - register_op_categories, -) +from TraceLens.PerfModel.op_categories import register_op_categories from TraceLens.PerfModel.torch_op_mapping import ( OP_CATEGORY_REGISTRY, categorize_torch_op, @@ -101,14 +98,6 @@ def test_dict_cat2names_keeps_sheet_compatibility_view(): assert "TokenPermuteMaskMap" not in dict_cat2names["MoE_comm"] -def test_dict_cat2names_dynamic_fallback_for_legacy_callers(): - local_sheet_view = {"SDPA": ["MyCustomAttentionBackward"]} - assert ( - category_from_sheet_view("MyCustomAttentionBackward", local_sheet_view) - == "SDPA_bwd" - ) - - def test_register_op_category_extension_updates_registry_only(): registry = {} From f5505fd774b8df7bb075a6f4f60a9faee470e061 Mon Sep 17 00:00:00 2001 From: Jassani Date: Fri, 8 May 2026 15:51:04 -0400 Subject: [PATCH 05/10] PerfModel: remove legacy category compatibility maps Derive legacy sheet membership from perf-model metadata so the registry remains the only categorization source of truth. Co-authored-by: Cursor --- TraceLens/PerfModel/__init__.py | 2 - TraceLens/PerfModel/extensions/__init__.py | 3 +- .../extensions/pseudo_ops_perf_utils.py | 27 - TraceLens/PerfModel/op_categories.py | 85 +-- TraceLens/PerfModel/perf_model.py | 2 + TraceLens/PerfModel/torch_op_mapping.py | 16 +- .../Reporting/generate_perf_report_pytorch.py | 23 +- .../generate_perf_report_pytorch_inference.py | 23 +- TraceLens/TreePerf/tree_perf.py | 2 - ...erf_report_with_fusion_and_shortkernels.py | 17 +- .../generate_perf_report_megatron_lm.py | 21 +- examples/tree_perf_example.ipynb | 567 +++++++++--------- tests/test_pseudo_ops_extension.py | 27 +- .../test_torch_op_categorization_registry.py | 24 +- 14 files changed, 378 insertions(+), 461 deletions(-) diff --git a/TraceLens/PerfModel/__init__.py b/TraceLens/PerfModel/__init__.py index e5558c2e8..48b37ef21 100644 --- a/TraceLens/PerfModel/__init__.py +++ b/TraceLens/PerfModel/__init__.py @@ -7,8 +7,6 @@ from .perf_model import * # Import everything from perf_model from .torch_op_mapping import ( op_to_perf_model_class_map, - dict_cat2names, - dict_base_class2category, ) __all__ = [name for name in dir() if not name.startswith("_")] diff --git a/TraceLens/PerfModel/extensions/__init__.py b/TraceLens/PerfModel/extensions/__init__.py index 230c5043e..cef661464 100644 --- a/TraceLens/PerfModel/extensions/__init__.py +++ b/TraceLens/PerfModel/extensions/__init__.py @@ -49,7 +49,7 @@ aiter_reduce_scatter, aiter_all_gather_reg, ) -from .pseudo_ops_perf_utils import get_pseudo_op_mappings, get_pseudo_op_categories +from .pseudo_ops_perf_utils import get_pseudo_op_mappings __all__ = [ # Base classes @@ -92,5 +92,4 @@ "custom_ar_qr_all_reduce", # Utility functions "get_pseudo_op_mappings", - "get_pseudo_op_categories", ] diff --git a/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py b/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py index f8df92637..ab80edc05 100644 --- a/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py +++ b/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py @@ -104,30 +104,3 @@ def get_pseudo_op_mappings(): return pseudo_op_mappings - -def get_pseudo_op_categories(): - """ - Return a dictionary mapping pseudo-op base classes to their performance categories. - - Returns: - dict: Mapping of base classes to category names - """ - - pseudo_op_categories = { - moe_perf_model_extensions.FusedMoE: "MoE_fused", - moe_perf_model_extensions.UnfusedMoE_Up: "MoE_unfused", - moe_perf_model_extensions.UnfusedMoE_Down: "MoE_unfused", - attention_perf_model_extensions.InferenceAttention: "InferenceAttention", - perf_model_extensions.gemm_a8w8_blockscale: "GEMM", - perf_model_extensions.batched_gemm_a16wfp4: "GEMM", - attention_perf_model_extensions.mha_varlen_fwd: "InferenceAttention", - perf_model_extensions.GroupQuant: "GroupQuant", - perf_model_extensions.gemm_a16w16_atomic_: "GEMM", - perf_model_extensions.batched_gemm_a8w8: "GEMM", - rmsnorm_perf_model_extensions.RMSNorm: "RMSNorm", - custom_collectives_perf_model_extensions.CustomCollective: "CustomCollective", - custom_collectives_perf_model_extensions.custom_ar_all_reduce: "CustomCollective", - perf_model_extensions.aiter_silu_and_mul: "UnaryElementwise", - } - - return pseudo_op_categories diff --git a/TraceLens/PerfModel/op_categories.py b/TraceLens/PerfModel/op_categories.py index 73571fe9d..7c5c9a464 100644 --- a/TraceLens/PerfModel/op_categories.py +++ b/TraceLens/PerfModel/op_categories.py @@ -9,16 +9,15 @@ Perf model classes declare their own output categories via ``category`` and, when linked-forward backward metrics are intentionally supported, ``bwd_category``. Category-only ops that do not have a perf model live in -``CATEGORY_ONLY_OP_CATEGORIES``. +``CATEGORY_ONLY_OP_MAPPING``. """ import re -from collections import defaultdict -from typing import DefaultDict, Dict, Iterable, List, Mapping, MutableMapping +from typing import Dict, Iterable, List, Mapping, MutableMapping from typing import Optional, Pattern, Tuple -CATEGORY_ONLY_OP_CATEGORIES: Dict[str, str] = { +CATEGORY_ONLY_OP_MAPPING: Dict[str, str] = { # CONV ops not present in op_to_perf_model_class_map. "aten::miopen_convolution": "CONV_fwd", "aten::cudnn_convolution": "CONV_fwd", @@ -76,14 +75,6 @@ ) -def _append_unique(target: List[str], names: Iterable[str]) -> None: - existing = set(target) - for name in names: - if name not in existing: - target.append(name) - existing.add(name) - - def _kernel_name_fallback(row) -> Optional[str]: kernel_details = row.get("kernel_details") if not kernel_details: @@ -97,14 +88,6 @@ def _kernel_name_fallback(row) -> Optional[str]: return None -def is_backward_op(op_name: str) -> bool: - return ( - op_name.endswith("_backward") - or op_name.endswith("Backward") - or op_name.endswith("_bwd") - ) - - def get_perf_model_category(perf_model_class: type, bwd: bool = False) -> Optional[str]: """Return the category declared by a perf model class.""" attr_name = "bwd_category" if bwd else "category" @@ -120,27 +103,10 @@ def get_perf_model_category(perf_model_class: type, bwd: bool = False) -> Option def sheet_category_from_final_category(category: str) -> str: """Return the legacy sheet family for a final categorization label.""" - category_to_sheet = { - "CONV_fwd": "CONV", - "CONV_bwd": "CONV", - "SDPA_fwd": "SDPA", - "SDPA_bwd": "SDPA", - "NORM_fwd": "Normalization", - "NORM_bwd": "Normalization", - "GroupedGEMM_fwd": "GroupedGEMM", - "GroupedGEMM_bwd": "GroupedGEMM", - "SSM_fwd": "SSM", - "SSM_bwd": "SSM", - "MoE_comm_fwd": "MoE_comm", - "MoE_comm_bwd": "MoE_comm", - "RoPE_fwd": "RoPE", - "RoPE_bwd": "RoPE", - "CrossEntropy_fwd": "CrossEntropy", - "CrossEntropy_bwd": "CrossEntropy", - "elementwise": "UnaryElementwise", - "reduce": "Reduce", - } - return category_to_sheet.get(category, category) + for suffix in ("_fwd", "_bwd"): + if category.endswith(suffix): + return category[: -len(suffix)] + return category def get_perf_model_sheet_category(perf_model_class: type) -> str: @@ -166,47 +132,24 @@ def build_op_category_registry( return registry -def build_dict_cat2names( +def build_sheet_category_to_op_names( op_to_perf_model_class_map: Mapping[str, type], -) -> DefaultDict[str, List[str]]: +) -> Dict[str, List[str]]: """Build the legacy ``category -> op names`` view used for report sheets.""" - dict_cat2names = defaultdict(list) # type: DefaultDict[str, List[str]] + sheet_category_to_op_names = {} # type: Dict[str, List[str]] for op_name, perf_model_class in op_to_perf_model_class_map.items(): - dict_cat2names[get_perf_model_sheet_category(perf_model_class)].append(op_name) - return dict_cat2names - - -def build_dict_base_class2category( - op_to_perf_model_class_map: Mapping[str, type], -) -> Dict[type, str]: - """Compatibility view for callers that still inspect base-class categories.""" - base_class2category: Dict[type, str] = {} - for perf_model_class in op_to_perf_model_class_map.values(): - base_classes = perf_model_class.__bases__ - if len(base_classes) != 1: - continue - base_class = base_classes[0] sheet_category = get_perf_model_sheet_category(perf_model_class) - existing = base_class2category.get(base_class) - if existing is not None and existing != sheet_category: - continue - base_class2category[base_class] = sheet_category - return base_class2category + sheet_category_to_op_names.setdefault(sheet_category, []).append(op_name) + return sheet_category_to_op_names def register_perf_model_categories( perf_model_extension: Mapping[str, type], registry: MutableMapping[str, str], - dict_cat2names: MutableMapping[str, List[str]], ) -> None: """Register categories for extension-provided perf models.""" for op_name, perf_model_class in perf_model_extension.items(): - category = get_perf_model_category(perf_model_class) - sheet_category = get_perf_model_sheet_category(perf_model_class) - registry[op_name] = category - if sheet_category not in dict_cat2names: - dict_cat2names[sheet_category] = [] - _append_unique(dict_cat2names[sheet_category], [op_name]) + registry[op_name] = get_perf_model_category(perf_model_class) def register_op_categories( @@ -217,7 +160,7 @@ def register_op_categories( registry.update(op_category_extension) -def categorize_torch_op_from_registry( +def _categorize_torch_op_from_registry( row, registry: Mapping[str, str], patterns: Iterable[Tuple[Pattern, str]] = OP_CATEGORY_PATTERNS, diff --git a/TraceLens/PerfModel/perf_model.py b/TraceLens/PerfModel/perf_model.py index a22949f89..ce867ee0e 100644 --- a/TraceLens/PerfModel/perf_model.py +++ b/TraceLens/PerfModel/perf_model.py @@ -3438,6 +3438,7 @@ class Reduce: category = "reduce" bwd_category = None + sheet_category = "Reduce" def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event @@ -4261,6 +4262,7 @@ def parse_list(input: str, dtype): class Normalization: category = "NORM_fwd" bwd_category = "NORM_bwd" + sheet_category = "Normalization" def __init__(self, event, arch=None, python_path=None, **kwargs): self.event = event diff --git a/TraceLens/PerfModel/torch_op_mapping.py b/TraceLens/PerfModel/torch_op_mapping.py index 9b1469c1a..162bfc230 100644 --- a/TraceLens/PerfModel/torch_op_mapping.py +++ b/TraceLens/PerfModel/torch_op_mapping.py @@ -7,11 +7,9 @@ from . import perf_model from .extensions import get_pseudo_op_mappings from .op_categories import ( - CATEGORY_ONLY_OP_CATEGORIES, - build_dict_base_class2category, - build_dict_cat2names, + CATEGORY_ONLY_OP_MAPPING, + _categorize_torch_op_from_registry, build_op_category_registry, - categorize_torch_op_from_registry, ) op_to_perf_model_class_map = { @@ -201,15 +199,9 @@ for op in reduce_ops: op_to_perf_model_class_map[op] = perf_model.aten_reduce -# Compatibility view for older callers that inspect base-class categories. New -# categorization reads ``category`` / ``bwd_category`` from perf model classes. -dict_base_class2category = build_dict_base_class2category(op_to_perf_model_class_map) - -dict_cat2names = build_dict_cat2names(op_to_perf_model_class_map) - OP_CATEGORY_REGISTRY = build_op_category_registry( op_to_perf_model_class_map, - category_only_ops=CATEGORY_ONLY_OP_CATEGORIES, + category_only_ops=CATEGORY_ONLY_OP_MAPPING, ) @@ -234,7 +226,7 @@ def categorize_torch_op(row): Note: Backward variants and auxiliary ops (TokenPermuteMaskMap, etc.) are categorization-only (timing without GFLOPS or TB/s). """ - return categorize_torch_op_from_registry( + return _categorize_torch_op_from_registry( row, OP_CATEGORY_REGISTRY, ) diff --git a/TraceLens/Reporting/generate_perf_report_pytorch.py b/TraceLens/Reporting/generate_perf_report_pytorch.py index bdb23dbcd..5d425a174 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch.py @@ -17,6 +17,7 @@ import pandas as pd from TraceLens import NcclAnalyser, TraceToTree, TraceDiff, TreePerfAnalyzer +from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names from TraceLens.Reporting.reporting_utils import request_install @@ -146,7 +147,6 @@ def apply_extension(perf_analyzer, extension_path): register_perf_model_categories( perf_model_extension, OP_CATEGORY_REGISTRY, - perf_analyzer.dict_cat2names, ) if hasattr(extension, "op_category_extension"): print(f"Applying op category extension from {extension_path}") @@ -378,16 +378,19 @@ def generate_perf_report_pytorch( ) # Dictionary to hold the op-specific DataFrames perf_metrics_dfs = {} + sheet_category_to_op_names = build_sheet_category_to_op_names( + perf_analyzer.op_to_perf_model_class_map + ) - for op_cat, op_names in perf_analyzer.dict_cat2names.items(): - # Filter events belonging to the current category + for sheet_category, op_names in sheet_category_to_op_names.items(): + # Filter events belonging to the current legacy sheet category op_events = [ event for event in perf_analyzer.tree.events if event["name"] in op_names ] - if op_cat in [ + if sheet_category in [ "GEMM", "UnaryElementwise", "BinaryElementwise", @@ -408,7 +411,7 @@ def generate_perf_report_pytorch( new_col_name="trunc_kernel_details", ) if not df_ops.empty: - perf_metrics_dfs[op_cat] = df_ops + perf_metrics_dfs[sheet_category] = df_ops if include_overlap_info: df_ops_overlapping_kernels = ( perf_analyzer.summarize_df_perf_metrics( @@ -429,7 +432,7 @@ def generate_perf_report_pytorch( new_col_name="trunc_overlapping_kernels_details", ) if not df_ops_overlapping_kernels.empty: - perf_metrics_dfs[f"{op_cat}_kl_overlap"] = ( + perf_metrics_dfs[f"{sheet_category}_kl_overlap"] = ( df_ops_overlapping_kernels ) else: @@ -497,9 +500,9 @@ def generate_perf_report_pytorch( ] df_ops_bwd = df_ops_bwd[~df_ops_bwd["name"].isin(fwd_op_names)] if not df_ops_fwd.empty: - perf_metrics_dfs[f"{op_cat}_fwd"] = df_ops_fwd + perf_metrics_dfs[f"{sheet_category}_fwd"] = df_ops_fwd if not df_ops_bwd.empty: - perf_metrics_dfs[f"{op_cat}_bwd"] = df_ops_bwd + perf_metrics_dfs[f"{sheet_category}_bwd"] = df_ops_bwd if include_overlap_info: df_ops_fwd_overlapping_kernels = ( @@ -576,11 +579,11 @@ def generate_perf_report_pytorch( ~df_ops_bwd_overlapping_kernels["name"].isin(fwd_op_names) ] if not df_ops_fwd_overlapping_kernels.empty: - perf_metrics_dfs[f"{op_cat}_fwd_kl_overlap"] = ( + perf_metrics_dfs[f"{sheet_category}_fwd_kl_overlap"] = ( df_ops_fwd_overlapping_kernels ) if not df_ops_bwd_overlapping_kernels.empty: - perf_metrics_dfs[f"{op_cat}_bwd_kl_overlap"] = ( + perf_metrics_dfs[f"{sheet_category}_bwd_kl_overlap"] = ( df_ops_bwd_overlapping_kernels ) diff --git a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py index a31e9acb6..c3d57e66b 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py @@ -23,6 +23,7 @@ import zipfile from TraceLens import NcclAnalyser, TraceToTree, TraceDiff, TreePerfAnalyzer +from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names from TraceLens.Reporting.reporting_utils import request_install from TraceLens.util import TraceEventUtils from TraceLens.Trace2Tree.trace_capture_merge_experimental import ( @@ -427,7 +428,6 @@ def apply_extension(perf_analyzer, extension_path): register_perf_model_categories( perf_model_extension, OP_CATEGORY_REGISTRY, - perf_analyzer.dict_cat2names, ) if hasattr(extension, "op_category_extension"): print(f"Applying op category extension from {extension_path}") @@ -687,8 +687,11 @@ def generate_perf_report_pytorch( ) # Dictionary to hold the op-specific DataFrames perf_metrics_dfs = {} - for op_cat, op_names in perf_analyzer.dict_cat2names.items(): - # Filter events belonging to the current category + sheet_category_to_op_names = build_sheet_category_to_op_names( + perf_analyzer.op_to_perf_model_class_map + ) + for sheet_category, op_names in sheet_category_to_op_names.items(): + # Filter events belonging to the current legacy sheet category op_events = [ event for event in perf_analyzer.tree.events @@ -697,7 +700,7 @@ def generate_perf_report_pytorch( if len(op_events) == 0: continue # Skip categories with no events - if op_cat in ["GEMM", "UnaryElementwise", "BinaryElementwise"]: + if sheet_category in ["GEMM", "UnaryElementwise", "BinaryElementwise"]: # For GEMM: create a single table that covers both fwd and bwd. df_ops_raw = perf_analyzer.build_df_perf_metrics( op_events, bwd=False, include_kernel_details=True, include_args=True @@ -713,7 +716,7 @@ def generate_perf_report_pytorch( new_col_name="trunc_kernel_details", ) if not df_ops.empty: - perf_metrics_dfs[op_cat] = df_ops + perf_metrics_dfs[sheet_category] = df_ops if include_overlap_info: df_ops_overlapping_kernels = ( perf_analyzer.summarize_df_perf_metrics( @@ -734,7 +737,7 @@ def generate_perf_report_pytorch( new_col_name="trunc_overlapping_kernels_details", ) if not df_ops_overlapping_kernels.empty: - perf_metrics_dfs[f"{op_cat}_kl_overlap"] = ( + perf_metrics_dfs[f"{sheet_category}_kl_overlap"] = ( df_ops_overlapping_kernels ) else: @@ -790,9 +793,9 @@ def generate_perf_report_pytorch( if filtered_df_bwd_ops is not None: df_ops_bwd = pd.concat([df_ops_bwd, filtered_df_bwd_ops]) if not df_ops_bwd.empty: - perf_metrics_dfs[f"{op_cat}_bwd"] = df_ops_bwd + perf_metrics_dfs[f"{sheet_category}_bwd"] = df_ops_bwd if not df_ops_fwd.empty: - perf_metrics_dfs[f"{op_cat}_fwd"] = df_ops_fwd + perf_metrics_dfs[f"{sheet_category}_fwd"] = df_ops_fwd if include_overlap_info: df_ops_fwd_overlapping_kernels = ( @@ -862,11 +865,11 @@ def generate_perf_report_pytorch( ] ) if not df_ops_bwd_overlapping_kernels.empty: - perf_metrics_dfs[f"{op_cat}_bwd_kl_overlap"] = ( + perf_metrics_dfs[f"{sheet_category}_bwd_kl_overlap"] = ( df_ops_bwd_overlapping_kernels ) if not df_ops_fwd_overlapping_kernels.empty: - perf_metrics_dfs[f"{op_cat}_fwd_kl_overlap"] = ( + perf_metrics_dfs[f"{sheet_category}_fwd_kl_overlap"] = ( df_ops_fwd_overlapping_kernels ) diff --git a/TraceLens/TreePerf/tree_perf.py b/TraceLens/TreePerf/tree_perf.py index 382237e6a..94316bae2 100644 --- a/TraceLens/TreePerf/tree_perf.py +++ b/TraceLens/TreePerf/tree_perf.py @@ -26,7 +26,6 @@ from ..PerfModel.jax_op_mapping import jax_op_to_perf_model_class_map from ..PerfModel.torch_op_mapping import ( categorize_torch_op, - dict_cat2names, op_to_perf_model_class_map, ) from ..PerfModel.op_categories import get_perf_model_category @@ -245,7 +244,6 @@ def __init__( self.op_to_perf_model_class_map = op_to_perf_model_class_map self.op_categorizer = categorize_torch_op - self.dict_cat2names = dict_cat2names def check_gpu_only(self): for event in self.tree.events: diff --git a/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py b/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py index 4c000ba88..510e52389 100644 --- a/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py +++ b/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py @@ -9,7 +9,7 @@ import pandas as pd import numpy as np from TraceLens import TreePerfAnalyzer -from TraceLens.PerfModel import dict_cat2names +from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names def get_next_host_op(perf_analyzer, host_op): @@ -328,13 +328,16 @@ def main(): # Roofline study op_dfs = {} - for op_cat, op_names in dict_cat2names.items(): - # Filter events belonging to the current category + sheet_category_to_op_names = build_sheet_category_to_op_names( + perf_analyzer.op_to_perf_model_class_map + ) + for sheet_category, op_names in sheet_category_to_op_names.items(): + # Filter events belonging to the current legacy sheet category op_events = [ event for event in perf_analyzer.tree.events if event["name"] in op_names ] - if op_cat in ["GEMM", "UnaryElementwise", "BinaryElementwise"]: + if sheet_category in ["GEMM", "UnaryElementwise", "BinaryElementwise"]: # For GEMM: create a single table that covers both fwd and bwd. df_ops = perf_analyzer.build_df_perf_metrics( op_events, bwd=False, include_kernel_names=True, include_args=True @@ -342,7 +345,7 @@ def main(): df_ops = perf_analyzer.summarize_df_perf_metrics(df_ops, agg_metrics) if args.topk_roofline_ops is not None: df_ops = df_ops.head(args.topk_roofline_ops) - op_dfs[op_cat] = df_ops + op_dfs[sheet_category] = df_ops else: # For FLASH_ATTN and CONV: create separate tables for forward and backward passes. df_ops_fwd = perf_analyzer.build_df_perf_metrics( @@ -361,8 +364,8 @@ def main(): ) if args.topk_roofline_ops is not None: df_ops_bwd = df_ops_bwd.head(args.topk_roofline_ops) - op_dfs[f"{op_cat}_fwd"] = df_ops_fwd - op_dfs[f"{op_cat}_bwd"] = df_ops_bwd + op_dfs[f"{sheet_category}_fwd"] = df_ops_fwd + op_dfs[f"{sheet_category}_bwd"] = df_ops_bwd # Write all DataFrames to separate sheets in an Excel workbook with pd.ExcelWriter(args.output_xlsx_path) as writer: diff --git a/examples/custom_workflows/generate_perf_report_megatron_lm.py b/examples/custom_workflows/generate_perf_report_megatron_lm.py index 245c14f21..288daf6fd 100644 --- a/examples/custom_workflows/generate_perf_report_megatron_lm.py +++ b/examples/custom_workflows/generate_perf_report_megatron_lm.py @@ -9,7 +9,7 @@ import pandas as pd from TraceLens import TraceToTree from TraceLens import TreePerfAnalyzer -from TraceLens.PerfModel import dict_cat2names +from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names from TraceLens.PerfModel import SDPA @@ -129,17 +129,20 @@ def main(): # Dictionary to hold the op-specific DataFrames op_dfs = {} - # update the dict_cat2names to include FusedAttnFunc - dict_cat2names["SDPA"].append("FusedAttnFunc") dict_name_to_custom_perf_model = {"FusedAttnFunc": transformer_engine_attention} + op_to_perf_model_class_map = dict(perf_analyzer.op_to_perf_model_class_map) + op_to_perf_model_class_map.update(dict_name_to_custom_perf_model) + sheet_category_to_op_names = build_sheet_category_to_op_names( + op_to_perf_model_class_map + ) - for op_cat, op_names in dict_cat2names.items(): - # Filter events belonging to the current category + for sheet_category, op_names in sheet_category_to_op_names.items(): + # Filter events belonging to the current legacy sheet category op_events = [ event for event in perf_analyzer.tree.events if event["name"] in op_names ] - if op_cat in ["GEMM", "UnaryElementwise", "BinaryElementwise"]: + if sheet_category in ["GEMM", "UnaryElementwise", "BinaryElementwise"]: # For GEMM: create a single table that covers both fwd and bwd. df_ops = perf_analyzer.build_df_perf_metrics( op_events, @@ -149,7 +152,7 @@ def main(): dict_name_to_perf_model=dict_name_to_custom_perf_model, ) df_ops = perf_analyzer.summarize_df_perf_metrics(df_ops, agg_metrics) - op_dfs[op_cat] = df_ops + op_dfs[sheet_category] = df_ops else: # For FLASH_ATTN and CONV: create separate tables for forward and backward passes. df_ops_fwd = perf_analyzer.build_df_perf_metrics( @@ -172,8 +175,8 @@ def main(): df_ops_bwd = perf_analyzer.summarize_df_perf_metrics( df_ops_bwd, agg_metrics ) - op_dfs[f"{op_cat}_fwd"] = df_ops_fwd - op_dfs[f"{op_cat}_bwd"] = df_ops_bwd + op_dfs[f"{sheet_category}_fwd"] = df_ops_fwd + op_dfs[f"{sheet_category}_bwd"] = df_ops_bwd # Write all DataFrames to separate sheets in an Excel workbook with pd.ExcelWriter(args.output_xlsx_path) as writer: diff --git a/examples/tree_perf_example.ipynb b/examples/tree_perf_example.ipynb index 7bda0df57..077fd8cd6 100755 --- a/examples/tree_perf_example.ipynb +++ b/examples/tree_perf_example.ipynb @@ -1,284 +1,287 @@ { - "cells": [ - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "" - ] + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "" + ] + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# This notebook demonstrates key features of the TraceLens library\n", + "# We encourage users to walk through the notebook and play with the code\n", + "# to get a feel for how the library works.\n", + "\n", + "# For production cases, we recommend using the TraceLens/examples/generate_perf_report.py script\n", + "\n", + "from pprint import pprint\n", + "import json\n", + "import pandas as pd\n", + "from TraceLens import TreePerfAnalyzer" + ], + "execution_count": null, + "outputs": [], + "id": "84eb5f1b-38e2-4b2d-9fcf-55c6ed7fb1dc" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# replace by your profile path, it can be a single rank profile from a multi gpu run as well\n", + "path = \"/path/to/profile.json\"\n", + "perf_analyzer = TreePerfAnalyzer.from_file(path)" + ], + "execution_count": null, + "outputs": [], + "id": "caf812d2-5b1d-4285-b9d6-8078173ecb27" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# get breakdown of gpu timeline - busy time, idle time, communication time, etc\n", + "perf_analyzer.get_df_gpu_timeline()" + ], + "execution_count": null, + "outputs": [], + "id": "0e7d4ac1-2cbe-49b8-97fa-d58eb0873b04" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# table of all lowest-level CPU operations (from the call stack perspective)\n", + "# and the time they \"induce\" on the GPU\n", + "df_kernel_launchers = perf_analyzer.get_df_kernel_launchers(include_kernel_details=True)\n", + "df_kernel_launchers.round(2).head()" + ], + "execution_count": null, + "outputs": [], + "id": "6f16ae0c-9d06-4215-8282-ba207af928fc" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# group by op name and summarize\n", + "# this gives an op wise breakdown of gpu time\n", + "df_kernel_launchers_summary = perf_analyzer.get_df_kernel_launchers_summary(\n", + " df_kernel_launchers\n", + ")\n", + "df_kernel_launchers_summary.round(2).head()" + ], + "execution_count": null, + "outputs": [], + "id": "71bf24ea-fb43-42cf-b09f-f5f464142d22" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Generate a detailed breakdown of unique argument combinations for all kernel-launching CPU ops.\n", + "# For each unique (op name + input dims/types/strides/concrete args), this groups and aggregates GPU time,\n", + "# helping identify which op and its arguments are the most time-consuming.\n", + "perf_analyzer.get_df_kernel_launchers_unique_args(df_kernel_launchers, include_pct=True)" + ], + "execution_count": null, + "outputs": [], + "id": "117537e4-ccc3-4c59-af9a-6cbe85283c40" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Same as above, but restricted to a specific op type — e.g., only `aten::mm`.\n", + "# Useful for drilling into the breakdown of a single op, such as mm, addmm, convolution, etc.\n", + "perf_analyzer.get_df_kernel_launchers_unique_args(\n", + " df_kernel_launchers, event_name=\"aten::mm\", include_pct=True\n", + ")" + ], + "execution_count": null, + "outputs": [], + "id": "03ce0ec9-725e-4982-ae85-99d8d83a6f36" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Roofline for ops\n", + "# currently we have GEMM, CONV fwd+bwd, FA\n", + "# many more coming soon\n", + "\n", + "# Example 1 GEMM\n", + "from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names\n", + "\n", + "sheet_category_to_op_names = build_sheet_category_to_op_names(\n", + " perf_analyzer.op_to_perf_model_class_map\n", + ")\n", + "gemm_op_names = sheet_category_to_op_names[\"GEMM\"]\n", + "gemm_events = [\n", + " event for event in perf_analyzer.tree.events if event[\"name\"] in gemm_op_names\n", + "]\n", + "print(f\"Found {len(gemm_events)} gemm events\")\n", + "\n", + "# take an example event and compute perf metrics\n", + "gemm_event = gemm_events[0]\n", + "print(\"Event dict:\")\n", + "pprint(gemm_event)\n", + "print(\"Perf metrics dict:\")\n", + "pprint(perf_analyzer.compute_perf_metrics(gemm_event))" + ], + "execution_count": null, + "outputs": [], + "id": "3a90c4f9-07bf-429f-840a-7e297542b4c3" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# build table for compute perf metrics for all gemm events\n", + "# include_kernel_details=True will add a column with the list of kernel name launched by the CPU op\n", + "df_gemm_ops = perf_analyzer.build_df_perf_metrics(\n", + " gemm_events, include_kernel_details=True\n", + ")\n", + "df_gemm_ops.head()" + ], + "execution_count": null, + "outputs": [], + "id": "f7a69bed-50fa-4475-bb41-34b2e1f8ee01" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# summarize by grouping across params M K N and bias and computing aggregate metrics\n", + "perf_analyzer.summarize_df_perf_metrics(df_gemm_ops, [\"mean\"])" + ], + "execution_count": null, + "outputs": [], + "id": "1bf969c8-2313-40f7-aa83-20568b7ac846" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Example 2a sdpa fwd\n", + "sdpa_op_names = sheet_category_to_op_names[\"SDPA\"]\n", + "sdpa_events = [\n", + " event for event in perf_analyzer.tree.events if event[\"name\"] in sdpa_op_names\n", + "]\n", + "df_sdpa_fwd_ops = perf_analyzer.build_df_perf_metrics(sdpa_events)\n", + "perf_analyzer.summarize_df_perf_metrics(df_sdpa_fwd_ops, [\"mean\"])" + ], + "execution_count": null, + "outputs": [], + "id": "c29741fc-4f49-441d-99d5-062396873b8a" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Example 2b sdpa bwd\n", + "# Note: bwd events for a fwd pass event are found\n", + "# by traversing the autograd links.\n", + "df_sdpa_bwd_ops = perf_analyzer.build_df_perf_metrics(sdpa_events, bwd=True)\n", + "perf_analyzer.summarize_df_perf_metrics(df_sdpa_bwd_ops, [\"mean\"])" + ], + "execution_count": null, + "outputs": [], + "id": "4dc2bba8-83c5-4e35-8571-be7ee1b6b1b0" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Example 3a conv fwd\n", + "conv_op_names = sheet_category_to_op_names[\"CONV\"]\n", + "conv_events = [\n", + " event for event in perf_analyzer.tree.events if event[\"name\"] in conv_op_names\n", + "]\n", + "df_conv_fwd_ops = perf_analyzer.build_df_perf_metrics(conv_events)\n", + "perf_analyzer.summarize_df_perf_metrics(df_conv_fwd_ops, [\"mean\"])" + ], + "execution_count": null, + "outputs": [], + "id": "0b606843-ea29-4286-b9ba-a180ea1b5534" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Example 3b conv bwd\n", + "df_conv_bwd_ops = perf_analyzer.build_df_perf_metrics(conv_events, bwd=True)\n", + "perf_analyzer.summarize_df_perf_metrics(df_conv_bwd_ops, [\"mean\"])" + ], + "execution_count": null, + "outputs": [], + "id": "e1d702d8-fbf4-45bb-82d1-ac3b5961ccd7" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Example 4 unary elementwise\n", + "\n", + "unary_elemwise_op_names = sheet_category_to_op_names[\"UnaryElementwise\"]\n", + "unary_elementwise_events = [\n", + " event\n", + " for event in perf_analyzer.tree.events\n", + " if event[\"name\"] in unary_elemwise_op_names\n", + "]\n", + "df_unary_elementwise_ops = perf_analyzer.build_df_perf_metrics(unary_elementwise_events)\n", + "perf_analyzer.summarize_df_perf_metrics(df_unary_elementwise_ops, [\"mean\"])" + ], + "execution_count": null, + "outputs": [], + "id": "ec56cc71-580c-4afe-b83e-da27148967b6" + }, + { + "cell_type": "code", + "metadata": {}, + "source": [ + "# Example 5 binary elementwise\n", + "binary_elemwise_op_names = sheet_category_to_op_names[\"BinaryElementwise\"]\n", + "binary_elementwise_events = [\n", + " event\n", + " for event in perf_analyzer.tree.events\n", + " if event[\"name\"] in binary_elemwise_op_names\n", + "]\n", + "df_binary_elementwise_ops = perf_analyzer.build_df_perf_metrics(\n", + " binary_elementwise_events\n", + ")\n", + "perf_analyzer.summarize_df_perf_metrics(df_binary_elementwise_ops, [\"mean\"])" + ], + "execution_count": null, + "outputs": [], + "id": "c521592f-9ded-487a-9694-d8307c438772" + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.10.16" + } }, - { - "cell_type": "code", - "execution_count": null, - "id": "84eb5f1b-38e2-4b2d-9fcf-55c6ed7fb1dc", - "metadata": {}, - "outputs": [], - "source": [ - "# This notebook demonstrates key features of the TraceLens library\n", - "# We encourage users to walk through the notebook and play with the code\n", - "# to get a feel for how the library works.\n", - "\n", - "# For production cases, we recommend using the TraceLens/examples/generate_perf_report.py script\n", - "\n", - "from pprint import pprint\n", - "import json\n", - "import pandas as pd\n", - "from TraceLens import TreePerfAnalyzer" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "caf812d2-5b1d-4285-b9d6-8078173ecb27", - "metadata": {}, - "outputs": [], - "source": [ - "# replace by your profile path, it can be a single rank profile from a multi gpu run as well\n", - "path = \"/path/to/profile.json\"\n", - "perf_analyzer = TreePerfAnalyzer.from_file(path)" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "0e7d4ac1-2cbe-49b8-97fa-d58eb0873b04", - "metadata": {}, - "outputs": [], - "source": [ - "# get breakdown of gpu timeline - busy time, idle time, communication time, etc\n", - "perf_analyzer.get_df_gpu_timeline()" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "6f16ae0c-9d06-4215-8282-ba207af928fc", - "metadata": {}, - "outputs": [], - "source": [ - "# table of all lowest-level CPU operations (from the call stack perspective)\n", - "# and the time they \"induce\" on the GPU\n", - "df_kernel_launchers = perf_analyzer.get_df_kernel_launchers(include_kernel_details=True)\n", - "df_kernel_launchers.round(2).head()" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "71bf24ea-fb43-42cf-b09f-f5f464142d22", - "metadata": {}, - "outputs": [], - "source": [ - "# group by op name and summarize\n", - "# this gives an op wise breakdown of gpu time\n", - "df_kernel_launchers_summary = perf_analyzer.get_df_kernel_launchers_summary(\n", - " df_kernel_launchers\n", - ")\n", - "df_kernel_launchers_summary.round(2).head()" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "117537e4-ccc3-4c59-af9a-6cbe85283c40", - "metadata": {}, - "outputs": [], - "source": [ - "# Generate a detailed breakdown of unique argument combinations for all kernel-launching CPU ops.\n", - "# For each unique (op name + input dims/types/strides/concrete args), this groups and aggregates GPU time,\n", - "# helping identify which op and its arguments are the most time-consuming.\n", - "perf_analyzer.get_df_kernel_launchers_unique_args(df_kernel_launchers, include_pct=True)" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "03ce0ec9-725e-4982-ae85-99d8d83a6f36", - "metadata": {}, - "outputs": [], - "source": [ - "# Same as above, but restricted to a specific op type — e.g., only `aten::mm`.\n", - "# Useful for drilling into the breakdown of a single op, such as mm, addmm, convolution, etc.\n", - "perf_analyzer.get_df_kernel_launchers_unique_args(\n", - " df_kernel_launchers, event_name=\"aten::mm\", include_pct=True\n", - ")" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "3a90c4f9-07bf-429f-840a-7e297542b4c3", - "metadata": {}, - "outputs": [], - "source": [ - "# Roofline for ops\n", - "# currently we have GEMM, CONV fwd+bwd, FA\n", - "# many more coming soon\n", - "\n", - "# Example 1 GEMM\n", - "from TraceLens.PerfModel import dict_cat2names\n", - "\n", - "gemm_op_names = dict_cat2names[\"GEMM\"]\n", - "gemm_events = [\n", - " event for event in perf_analyzer.tree.events if event[\"name\"] in gemm_op_names\n", - "]\n", - "print(f\"Found {len(gemm_events)} gemm events\")\n", - "\n", - "# take an example event and compute perf metrics\n", - "gemm_event = gemm_events[0]\n", - "print(\"Event dict:\")\n", - "pprint(gemm_event)\n", - "print(\"Perf metrics dict:\")\n", - "pprint(perf_analyzer.compute_perf_metrics(gemm_event))" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "f7a69bed-50fa-4475-bb41-34b2e1f8ee01", - "metadata": {}, - "outputs": [], - "source": [ - "# build table for compute perf metrics for all gemm events\n", - "# include_kernel_details=True will add a column with the list of kernel name launched by the CPU op\n", - "df_gemm_ops = perf_analyzer.build_df_perf_metrics(\n", - " gemm_events, include_kernel_details=True\n", - ")\n", - "df_gemm_ops.head()" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "1bf969c8-2313-40f7-aa83-20568b7ac846", - "metadata": {}, - "outputs": [], - "source": [ - "# summarize by grouping across params M K N and bias and computing aggregate metrics\n", - "perf_analyzer.summarize_df_perf_metrics(df_gemm_ops, [\"mean\"])" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "c29741fc-4f49-441d-99d5-062396873b8a", - "metadata": {}, - "outputs": [], - "source": [ - "# Example 2a sdpa fwd\n", - "sdpa_op_names = dict_cat2names[\"SDPA\"]\n", - "sdpa_events = [\n", - " event for event in perf_analyzer.tree.events if event[\"name\"] in sdpa_op_names\n", - "]\n", - "df_sdpa_fwd_ops = perf_analyzer.build_df_perf_metrics(sdpa_events)\n", - "perf_analyzer.summarize_df_perf_metrics(df_sdpa_fwd_ops, [\"mean\"])" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "4dc2bba8-83c5-4e35-8571-be7ee1b6b1b0", - "metadata": {}, - "outputs": [], - "source": [ - "# Example 2b sdpa bwd\n", - "# Note: bwd events for a fwd pass event are found\n", - "# by traversing the autograd links.\n", - "df_sdpa_bwd_ops = perf_analyzer.build_df_perf_metrics(sdpa_events, bwd=True)\n", - "perf_analyzer.summarize_df_perf_metrics(df_sdpa_bwd_ops, [\"mean\"])" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "0b606843-ea29-4286-b9ba-a180ea1b5534", - "metadata": {}, - "outputs": [], - "source": [ - "# Example 3a conv fwd\n", - "conv_op_names = dict_cat2names[\"CONV\"]\n", - "conv_events = [\n", - " event for event in perf_analyzer.tree.events if event[\"name\"] in conv_op_names\n", - "]\n", - "df_conv_fwd_ops = perf_analyzer.build_df_perf_metrics(conv_events)\n", - "perf_analyzer.summarize_df_perf_metrics(df_conv_fwd_ops, [\"mean\"])" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "e1d702d8-fbf4-45bb-82d1-ac3b5961ccd7", - "metadata": {}, - "outputs": [], - "source": [ - "# Example 3b conv bwd\n", - "df_conv_bwd_ops = perf_analyzer.build_df_perf_metrics(conv_events, bwd=True)\n", - "perf_analyzer.summarize_df_perf_metrics(df_conv_bwd_ops, [\"mean\"])" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "ec56cc71-580c-4afe-b83e-da27148967b6", - "metadata": {}, - "outputs": [], - "source": [ - "# Example 4 unary elementwise\n", - "\n", - "unary_elemwise_op_names = dict_cat2names[\"UnaryElementwise\"]\n", - "unary_elementwise_events = [\n", - " event\n", - " for event in perf_analyzer.tree.events\n", - " if event[\"name\"] in unary_elemwise_op_names\n", - "]\n", - "df_unary_elementwise_ops = perf_analyzer.build_df_perf_metrics(unary_elementwise_events)\n", - "perf_analyzer.summarize_df_perf_metrics(df_unary_elementwise_ops, [\"mean\"])" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "c521592f-9ded-487a-9694-d8307c438772", - "metadata": {}, - "outputs": [], - "source": [ - "# Example 5 binary elementwise\n", - "binary_elemwise_op_names = dict_cat2names[\"BinaryElementwise\"]\n", - "binary_elementwise_events = [\n", - " event\n", - " for event in perf_analyzer.tree.events\n", - " if event[\"name\"] in binary_elemwise_op_names\n", - "]\n", - "df_binary_elementwise_ops = perf_analyzer.build_df_perf_metrics(\n", - " binary_elementwise_events\n", - ")\n", - "perf_analyzer.summarize_df_perf_metrics(df_binary_elementwise_ops, [\"mean\"])" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": "Python 3 (ipykernel)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.10.16" - } - }, - "nbformat": 4, - "nbformat_minor": 5 -} + "nbformat": 4, + "nbformat_minor": 5 +} \ No newline at end of file diff --git a/tests/test_pseudo_ops_extension.py b/tests/test_pseudo_ops_extension.py index 5062cf206..9d6c183c4 100644 --- a/tests/test_pseudo_ops_extension.py +++ b/tests/test_pseudo_ops_extension.py @@ -554,20 +554,15 @@ def test_categorize_as_sdpa_bwd(self): def test_fused_attn_fwd_still_sdpa_fwd(self): """FusedAttnFunc (forward) must still be SDPA_fwd.""" from TraceLens.PerfModel.op_categories import register_perf_model_categories - from TraceLens.PerfModel.torch_op_mapping import ( - OP_CATEGORY_REGISTRY, - categorize_torch_op, - dict_cat2names, - ) + from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY + registry = dict(OP_CATEGORY_REGISTRY) register_perf_model_categories( {"FusedAttnFunc": perf_model_extension["FusedAttnFunc"]}, - OP_CATEGORY_REGISTRY, - dict_cat2names, + registry, ) - result = categorize_torch_op({"name": "FusedAttnFunc"}) - assert result == "SDPA_fwd", f"Expected SDPA_fwd, got {result}" + assert registry["FusedAttnFunc"] == "SDPA_fwd" class TestLayerNormFnPerfModel: @@ -648,23 +643,19 @@ def test_categorization_normalization(self): def test_categorize_as_norm_fwd_bwd(self): """Core categorizer must return NORM_fwd and NORM_bwd.""" from TraceLens.PerfModel.op_categories import register_perf_model_categories - from TraceLens.PerfModel.torch_op_mapping import ( - OP_CATEGORY_REGISTRY, - categorize_torch_op, - dict_cat2names, - ) + from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY + registry = dict(OP_CATEGORY_REGISTRY) register_perf_model_categories( { "LayerNormFn": te_layer_norm_fwd, "LayerNormFnBackward": te_layer_norm_bwd, }, - OP_CATEGORY_REGISTRY, - dict_cat2names, + registry, ) - assert categorize_torch_op({"name": "LayerNormFn"}) == "NORM_fwd" - assert categorize_torch_op({"name": "LayerNormFnBackward"}) == "NORM_bwd" + assert registry["LayerNormFn"] == "NORM_fwd" + assert registry["LayerNormFnBackward"] == "NORM_bwd" def test_perf_model_extension_registration(self): """LayerNormFn/LayerNormFnBackward must be registered in perf_model_extension.""" diff --git a/tests/test_torch_op_categorization_registry.py b/tests/test_torch_op_categorization_registry.py index fa11838d6..928089ccb 100644 --- a/tests/test_torch_op_categorization_registry.py +++ b/tests/test_torch_op_categorization_registry.py @@ -10,11 +10,13 @@ import pytest -from TraceLens.PerfModel.op_categories import register_op_categories +from TraceLens.PerfModel.op_categories import ( + build_sheet_category_to_op_names, + register_op_categories, +) from TraceLens.PerfModel.torch_op_mapping import ( OP_CATEGORY_REGISTRY, categorize_torch_op, - dict_cat2names, op_to_perf_model_class_map, ) @@ -89,13 +91,17 @@ def test_registry_covers_every_perf_model_op(): assert missing == [] -def test_dict_cat2names_keeps_sheet_compatibility_view(): - assert "aten::mm" in dict_cat2names["GEMM"] - assert "primus_turbo::grouped_gemm" in dict_cat2names["GroupedGEMM"] - assert "aten::convolution" in dict_cat2names["CONV"] - assert "aten::convolution_backward" in dict_cat2names["CONV"] - assert "FlashAttnFuncBackward" not in dict_cat2names["SDPA"] - assert "TokenPermuteMaskMap" not in dict_cat2names["MoE_comm"] +def test_sheet_category_to_op_names_keeps_legacy_sheet_view(): + sheet_category_to_op_names = build_sheet_category_to_op_names( + op_to_perf_model_class_map + ) + + assert "aten::mm" in sheet_category_to_op_names["GEMM"] + assert "primus_turbo::grouped_gemm" in sheet_category_to_op_names["GroupedGEMM"] + assert "aten::convolution" in sheet_category_to_op_names["CONV"] + assert "aten::convolution_backward" in sheet_category_to_op_names["CONV"] + assert "FlashAttnFuncBackward" not in sheet_category_to_op_names["SDPA"] + assert "TokenPermuteMaskMap" not in sheet_category_to_op_names["MoE_comm"] def test_register_op_category_extension_updates_registry_only(): From fd61967b286548ad11a820d1a4799e2afbd6706a Mon Sep 17 00:00:00 2001 From: Jassani Date: Fri, 8 May 2026 16:21:00 -0400 Subject: [PATCH 06/10] PerfModel: merge category helpers into torch mapping Keep torch op categorization in one module now that legacy compatibility maps are gone, and preserve RMSNorm as its own legacy report sheet. Co-authored-by: Cursor --- .../extensions/pseudo_ops_perf_utils.py | 1 - .../rmsnorm_perf_model_extensions.py | 1 + TraceLens/PerfModel/op_categories.py | 183 ------------------ TraceLens/PerfModel/torch_op_mapping.py | 175 ++++++++++++++++- .../Reporting/generate_perf_report_pytorch.py | 6 +- .../generate_perf_report_pytorch_inference.py | 6 +- TraceLens/TreePerf/tree_perf.py | 6 +- ...erf_report_with_fusion_and_shortkernels.py | 2 +- .../generate_perf_report_megatron_lm.py | 2 +- examples/tree_perf_example.ipynb | 2 +- tests/test_pseudo_ops_extension.py | 12 +- .../test_torch_op_categorization_registry.py | 10 +- 12 files changed, 196 insertions(+), 210 deletions(-) delete mode 100644 TraceLens/PerfModel/op_categories.py diff --git a/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py b/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py index ab80edc05..b2b944b2b 100644 --- a/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py +++ b/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py @@ -103,4 +103,3 @@ def get_pseudo_op_mappings(): } return pseudo_op_mappings - diff --git a/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py b/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py index 89e2ec4d0..c950a754b 100644 --- a/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/rmsnorm_perf_model_extensions.py @@ -16,6 +16,7 @@ class RMSNorm(CoreRMSNorm): category = "RMSNorm" bwd_category = None + sheet_category = "RMSNorm" class aiter_rms_norm(RMSNorm): diff --git a/TraceLens/PerfModel/op_categories.py b/TraceLens/PerfModel/op_categories.py deleted file mode 100644 index 7c5c9a464..000000000 --- a/TraceLens/PerfModel/op_categories.py +++ /dev/null @@ -1,183 +0,0 @@ -############################################################################### -# Copyright (c) 2024 - 2025 Advanced Micro Devices, Inc. All rights reserved. -# -# See LICENSE for license information. -############################################################################### - -"""Registry-based categorization of CPU torch ops. - -Perf model classes declare their own output categories via ``category`` and, -when linked-forward backward metrics are intentionally supported, -``bwd_category``. Category-only ops that do not have a perf model live in -``CATEGORY_ONLY_OP_MAPPING``. -""" - -import re -from typing import Dict, Iterable, List, Mapping, MutableMapping -from typing import Optional, Pattern, Tuple - - -CATEGORY_ONLY_OP_MAPPING: Dict[str, str] = { - # CONV ops not present in op_to_perf_model_class_map. - "aten::miopen_convolution": "CONV_fwd", - "aten::cudnn_convolution": "CONV_fwd", - # SDPA backward ops without direct perf models in core TraceLens. - "FlashAttnFuncBackward": "SDPA_bwd", - "FusedAttnFuncBackward": "SDPA_bwd", - "aten::_scaled_dot_product_cudnn_attention_backward": "SDPA_bwd", - "aten::_scaled_dot_product_efficient_attention_backward": "SDPA_bwd", - "aten::_scaled_dot_product_flash_attention_backward": "SDPA_bwd", - # SSM / Mamba category-only backward ops. - "MambaSplitConv1dScanCombinedFnBackward": "SSM_bwd", - "DaoAILab::_causal_conv1d_bwd_cpp": "SSM_bwd", - # MoE communication category-only ops. - "TokenPermuteMaskMap": "MoE_comm_fwd", - # Observed in MoE token-routing traces; tracked separately because the name - # itself is generic and may not always imply MoE communication. - "_OperationFuserAutogradFunction": "MoE_comm_fwd", - "MoEDispatchBackward": "MoE_comm_bwd", - "MoECombineBackward": "MoE_comm_bwd", - "TokenPermuteMaskMapBackward": "MoE_comm_bwd", - "_OperationFuserAutogradFunctionBackward": "MoE_comm_bwd", - # RoPE / CrossEntropy category-only backward ops. - "FusedRoPEFuncBackward": "RoPE_bwd", - "CrossEntropyFunctionBackward": "CrossEntropy_bwd", - # MoE auxiliary ops. - "aiter::moe_sorting_fwd": "MoE_aux", - "aiter::moe_sorting_opus_fwd": "MoE_aux", - "aiter::moe_align_block_size": "MoE_aux", - "_moe_C::moe_align_block_size": "MoE_aux", - "aiter::fused_moe_->_fused_dynamic_mxfp4_quant_moe_sort_kernel (Synthetic Op)": "MoE_aux", - "aiter::moe_sum": "MoE_aux", - "aiter::topk_softmax": "MoE_aux", - "aiter::topk_softmax_asm": "MoE_aux", - "aiter::topk_sigmoid": "MoE_aux", - "aiter::biased_grouped_topk_hip": "MoE_aux", - "aiter::grouped_topk": "MoE_aux", - "aiter::moe_fused_gate": "MoE_aux", - # InferenceAttention extras (KV-cache writes). - "_C_cache_ops::reshape_and_cache_flash": "InferenceAttention", - "_C_cache_ops::concat_and_cache_mla": "InferenceAttention", -} - - -OP_CATEGORY_PATTERNS: List[Tuple[Pattern, str]] = [ - (re.compile(r"^triton"), "triton"), - (re.compile(r"^record_param_comms"), "record_param_comms"), -] - - -_KERNEL_NAME_PREFIX = "void at::native" -_KERNEL_NAME_FALLBACK_RULES: Tuple[Tuple[str, str], ...] = ( - ("elementwise", "elementwise"), - ("reduce", "reduce"), - ("multi_tensor_apply", "multi_tensor_apply"), -) - - -def _kernel_name_fallback(row) -> Optional[str]: - kernel_details = row.get("kernel_details") - if not kernel_details: - return None - kernel_name = kernel_details[0].get("name", "") - if not kernel_name.startswith(_KERNEL_NAME_PREFIX): - return None - for needle, category in _KERNEL_NAME_FALLBACK_RULES: - if needle in kernel_name: - return category - return None - - -def get_perf_model_category(perf_model_class: type, bwd: bool = False) -> Optional[str]: - """Return the category declared by a perf model class.""" - attr_name = "bwd_category" if bwd else "category" - category = getattr(perf_model_class, attr_name, None) - if bwd: - return category - if category is None: - raise ValueError( - f"perf_model_class {perf_model_class} must define a category attribute" - ) - return category - - -def sheet_category_from_final_category(category: str) -> str: - """Return the legacy sheet family for a final categorization label.""" - for suffix in ("_fwd", "_bwd"): - if category.endswith(suffix): - return category[: -len(suffix)] - return category - - -def get_perf_model_sheet_category(perf_model_class: type) -> str: - """Return the legacy sheet category for a perf model class.""" - sheet_category = getattr(perf_model_class, "sheet_category", None) - if sheet_category is not None: - return sheet_category - return sheet_category_from_final_category(get_perf_model_category(perf_model_class)) - - -def build_op_category_registry( - op_to_perf_model_class_map: Mapping[str, type], - category_only_ops: Optional[Mapping[str, str]] = None, -) -> Dict[str, str]: - """Construct the flat ``op_name -> final category`` registry.""" - registry: Dict[str, str] = {} - for op_name, perf_model_class in op_to_perf_model_class_map.items(): - registry[op_name] = get_perf_model_category(perf_model_class) - - if category_only_ops: - registry.update(category_only_ops) - - return registry - - -def build_sheet_category_to_op_names( - op_to_perf_model_class_map: Mapping[str, type], -) -> Dict[str, List[str]]: - """Build the legacy ``category -> op names`` view used for report sheets.""" - sheet_category_to_op_names = {} # type: Dict[str, List[str]] - for op_name, perf_model_class in op_to_perf_model_class_map.items(): - sheet_category = get_perf_model_sheet_category(perf_model_class) - sheet_category_to_op_names.setdefault(sheet_category, []).append(op_name) - return sheet_category_to_op_names - - -def register_perf_model_categories( - perf_model_extension: Mapping[str, type], - registry: MutableMapping[str, str], -) -> None: - """Register categories for extension-provided perf models.""" - for op_name, perf_model_class in perf_model_extension.items(): - registry[op_name] = get_perf_model_category(perf_model_class) - - -def register_op_categories( - op_category_extension: Mapping[str, str], - registry: MutableMapping[str, str], -) -> None: - """Register explicit category-only op labels.""" - registry.update(op_category_extension) - - -def _categorize_torch_op_from_registry( - row, - registry: Mapping[str, str], - patterns: Iterable[Tuple[Pattern, str]] = OP_CATEGORY_PATTERNS, -) -> str: - """Return the category for ``row`` using explicit registry data.""" - name = row["name"] - - category = registry.get(name) - if category is not None: - return category - - for pattern, category in patterns: - if pattern.match(name): - return category - - fallback = _kernel_name_fallback(row) - if fallback is not None: - return fallback - - return "other" diff --git a/TraceLens/PerfModel/torch_op_mapping.py b/TraceLens/PerfModel/torch_op_mapping.py index 162bfc230..fec9dd786 100644 --- a/TraceLens/PerfModel/torch_op_mapping.py +++ b/TraceLens/PerfModel/torch_op_mapping.py @@ -4,14 +4,181 @@ # See LICENSE for license information. ############################################################################### +"""Torch op-name mappings and categorization helpers.""" + +import re +from typing import Dict, Iterable, List, Mapping, MutableMapping +from typing import Optional, Pattern, Tuple + from . import perf_model from .extensions import get_pseudo_op_mappings -from .op_categories import ( - CATEGORY_ONLY_OP_MAPPING, - _categorize_torch_op_from_registry, - build_op_category_registry, + +CATEGORY_ONLY_OP_MAPPING: Dict[str, str] = { + # CONV ops not present in op_to_perf_model_class_map. + "aten::miopen_convolution": "CONV_fwd", + "aten::cudnn_convolution": "CONV_fwd", + # SDPA backward ops without direct perf models in core TraceLens. + "FlashAttnFuncBackward": "SDPA_bwd", + "FusedAttnFuncBackward": "SDPA_bwd", + "aten::_scaled_dot_product_cudnn_attention_backward": "SDPA_bwd", + "aten::_scaled_dot_product_efficient_attention_backward": "SDPA_bwd", + "aten::_scaled_dot_product_flash_attention_backward": "SDPA_bwd", + # SSM / Mamba category-only backward ops. + "MambaSplitConv1dScanCombinedFnBackward": "SSM_bwd", + "DaoAILab::_causal_conv1d_bwd_cpp": "SSM_bwd", + # MoE communication category-only ops. + "TokenPermuteMaskMap": "MoE_comm_fwd", + # Observed in MoE token-routing traces; tracked separately because the name + # itself is generic and may not always imply MoE communication. + "_OperationFuserAutogradFunction": "MoE_comm_fwd", + "MoEDispatchBackward": "MoE_comm_bwd", + "MoECombineBackward": "MoE_comm_bwd", + "TokenPermuteMaskMapBackward": "MoE_comm_bwd", + "_OperationFuserAutogradFunctionBackward": "MoE_comm_bwd", + # RoPE / CrossEntropy category-only backward ops. + "FusedRoPEFuncBackward": "RoPE_bwd", + "CrossEntropyFunctionBackward": "CrossEntropy_bwd", + # MoE auxiliary ops. + "aiter::moe_sorting_fwd": "MoE_aux", + "aiter::moe_sorting_opus_fwd": "MoE_aux", + "aiter::moe_align_block_size": "MoE_aux", + "_moe_C::moe_align_block_size": "MoE_aux", + "aiter::fused_moe_->_fused_dynamic_mxfp4_quant_moe_sort_kernel (Synthetic Op)": "MoE_aux", + "aiter::moe_sum": "MoE_aux", + "aiter::topk_softmax": "MoE_aux", + "aiter::topk_softmax_asm": "MoE_aux", + "aiter::topk_sigmoid": "MoE_aux", + "aiter::biased_grouped_topk_hip": "MoE_aux", + "aiter::grouped_topk": "MoE_aux", + "aiter::moe_fused_gate": "MoE_aux", + # InferenceAttention extras (KV-cache writes). + "_C_cache_ops::reshape_and_cache_flash": "InferenceAttention", + "_C_cache_ops::concat_and_cache_mla": "InferenceAttention", +} + + +OP_CATEGORY_PATTERNS: List[Tuple[Pattern, str]] = [ + (re.compile(r"^triton"), "triton"), + (re.compile(r"^record_param_comms"), "record_param_comms"), +] + + +_KERNEL_NAME_PREFIX = "void at::native" +_KERNEL_NAME_FALLBACK_RULES: Tuple[Tuple[str, str], ...] = ( + ("elementwise", "elementwise"), + ("reduce", "reduce"), + ("multi_tensor_apply", "multi_tensor_apply"), ) + +def _kernel_name_fallback(row) -> Optional[str]: + kernel_details = row.get("kernel_details") + if not kernel_details: + return None + kernel_name = kernel_details[0].get("name", "") + if not kernel_name.startswith(_KERNEL_NAME_PREFIX): + return None + for needle, category in _KERNEL_NAME_FALLBACK_RULES: + if needle in kernel_name: + return category + return None + + +def get_perf_model_category(perf_model_class: type, bwd: bool = False) -> Optional[str]: + """Return the category declared by a perf model class.""" + attr_name = "bwd_category" if bwd else "category" + category = getattr(perf_model_class, attr_name, None) + if bwd: + return category + if category is None: + raise ValueError( + f"perf_model_class {perf_model_class} must define a category attribute" + ) + return category + + +def sheet_category_from_final_category(category: str) -> str: + """Return the legacy sheet family for a final categorization label.""" + for suffix in ("_fwd", "_bwd"): + if category.endswith(suffix): + return category[: -len(suffix)] + return category + + +def get_perf_model_sheet_category(perf_model_class: type) -> str: + """Return the legacy sheet category for a perf model class.""" + sheet_category = getattr(perf_model_class, "sheet_category", None) + if sheet_category is not None: + return sheet_category + return sheet_category_from_final_category(get_perf_model_category(perf_model_class)) + + +def build_op_category_registry( + op_to_perf_model_class_map: Mapping[str, type], + category_only_ops: Optional[Mapping[str, str]] = None, +) -> Dict[str, str]: + """Construct the flat ``op_name -> final category`` registry.""" + registry: Dict[str, str] = {} + for op_name, perf_model_class in op_to_perf_model_class_map.items(): + registry[op_name] = get_perf_model_category(perf_model_class) + + if category_only_ops: + registry.update(category_only_ops) + + return registry + + +def build_sheet_category_to_op_names( + op_to_perf_model_class_map: Mapping[str, type], +) -> Dict[str, List[str]]: + """Build the legacy ``category -> op names`` view used for report sheets.""" + sheet_category_to_op_names = {} # type: Dict[str, List[str]] + for op_name, perf_model_class in op_to_perf_model_class_map.items(): + sheet_category = get_perf_model_sheet_category(perf_model_class) + sheet_category_to_op_names.setdefault(sheet_category, []).append(op_name) + return sheet_category_to_op_names + + +def register_perf_model_categories( + perf_model_extension: Mapping[str, type], + registry: MutableMapping[str, str], +) -> None: + """Register categories for extension-provided perf models.""" + for op_name, perf_model_class in perf_model_extension.items(): + registry[op_name] = get_perf_model_category(perf_model_class) + + +def register_op_categories( + op_category_extension: Mapping[str, str], + registry: MutableMapping[str, str], +) -> None: + """Register explicit category-only op labels.""" + registry.update(op_category_extension) + + +def _categorize_torch_op_from_registry( + row, + registry: Mapping[str, str], + patterns: Iterable[Tuple[Pattern, str]] = OP_CATEGORY_PATTERNS, +) -> str: + """Return the category for ``row`` using explicit registry data.""" + name = row["name"] + + category = registry.get(name) + if category is not None: + return category + + for pattern, category in patterns: + if pattern.match(name): + return category + + fallback = _kernel_name_fallback(row) + if fallback is not None: + return fallback + + return "other" + + op_to_perf_model_class_map = { "aten::mm": perf_model.aten_mm, "aten::addmm": perf_model.aten_addmm, diff --git a/TraceLens/Reporting/generate_perf_report_pytorch.py b/TraceLens/Reporting/generate_perf_report_pytorch.py index 5d425a174..7e41a9aa1 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch.py @@ -17,7 +17,7 @@ import pandas as pd from TraceLens import NcclAnalyser, TraceToTree, TraceDiff, TreePerfAnalyzer -from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names +from TraceLens.PerfModel.torch_op_mapping import build_sheet_category_to_op_names from TraceLens.Reporting.reporting_utils import request_install @@ -120,11 +120,11 @@ def apply_extension(perf_analyzer, extension_path): extension_path = os.path.abspath(extension_path) extension_name = os.path.splitext(os.path.basename(extension_path))[0] - from TraceLens.PerfModel.op_categories import ( + from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, register_op_categories, register_perf_model_categories, ) - from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY spec = importlib.util.spec_from_file_location(extension_name, extension_path) extension = importlib.util.module_from_spec(spec) diff --git a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py index c3d57e66b..bc28aca90 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py @@ -23,7 +23,7 @@ import zipfile from TraceLens import NcclAnalyser, TraceToTree, TraceDiff, TreePerfAnalyzer -from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names +from TraceLens.PerfModel.torch_op_mapping import build_sheet_category_to_op_names from TraceLens.Reporting.reporting_utils import request_install from TraceLens.util import TraceEventUtils from TraceLens.Trace2Tree.trace_capture_merge_experimental import ( @@ -401,11 +401,11 @@ def apply_extension(perf_analyzer, extension_path): extension_path = os.path.abspath(extension_path) extension_name = os.path.splitext(os.path.basename(extension_path))[0] - from TraceLens.PerfModel.op_categories import ( + from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, register_op_categories, register_perf_model_categories, ) - from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY spec = importlib.util.spec_from_file_location(extension_name, extension_path) extension = importlib.util.module_from_spec(spec) diff --git a/TraceLens/TreePerf/tree_perf.py b/TraceLens/TreePerf/tree_perf.py index 94316bae2..a838d845c 100644 --- a/TraceLens/TreePerf/tree_perf.py +++ b/TraceLens/TreePerf/tree_perf.py @@ -26,9 +26,9 @@ from ..PerfModel.jax_op_mapping import jax_op_to_perf_model_class_map from ..PerfModel.torch_op_mapping import ( categorize_torch_op, + get_perf_model_category, op_to_perf_model_class_map, ) -from ..PerfModel.op_categories import get_perf_model_category from ..Trace2Tree.extensions import apply_pseudo_op_extensions from ..Trace2Tree.trace_capture_merge_experimental import merge_capture_trace_into_graph from ..Trace2Tree.trace_to_tree import JaxTraceToTree, TraceToTree @@ -1976,9 +1976,7 @@ def list_to_tuple(obj): "metrics_source": ( "own_perf_model" if has_own_perf_model - else "linked_forward_bwd" - if is_sole_bwd - else "kernel_time_only" + else "linked_forward_bwd" if is_sole_bwd else "kernel_time_only" ), "overlapping_kernel_names": event.get("overlapping_kernel_names"), "overlapping_kernels_details": event.get("overlapping_kernels_details"), diff --git a/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py b/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py index 510e52389..ebb4396a4 100644 --- a/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py +++ b/examples/archived/generate_perf_report_with_fusion_and_shortkernels.py @@ -9,7 +9,7 @@ import pandas as pd import numpy as np from TraceLens import TreePerfAnalyzer -from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names +from TraceLens.PerfModel.torch_op_mapping import build_sheet_category_to_op_names def get_next_host_op(perf_analyzer, host_op): diff --git a/examples/custom_workflows/generate_perf_report_megatron_lm.py b/examples/custom_workflows/generate_perf_report_megatron_lm.py index 288daf6fd..09ffb2683 100644 --- a/examples/custom_workflows/generate_perf_report_megatron_lm.py +++ b/examples/custom_workflows/generate_perf_report_megatron_lm.py @@ -9,7 +9,7 @@ import pandas as pd from TraceLens import TraceToTree from TraceLens import TreePerfAnalyzer -from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names +from TraceLens.PerfModel.torch_op_mapping import build_sheet_category_to_op_names from TraceLens.PerfModel import SDPA diff --git a/examples/tree_perf_example.ipynb b/examples/tree_perf_example.ipynb index 077fd8cd6..7a23e2f07 100755 --- a/examples/tree_perf_example.ipynb +++ b/examples/tree_perf_example.ipynb @@ -117,7 +117,7 @@ "# many more coming soon\n", "\n", "# Example 1 GEMM\n", - "from TraceLens.PerfModel.op_categories import build_sheet_category_to_op_names\n", + "from TraceLens.PerfModel.torch_op_mapping import build_sheet_category_to_op_names\n", "\n", "sheet_category_to_op_names = build_sheet_category_to_op_names(\n", " perf_analyzer.op_to_perf_model_class_map\n", diff --git a/tests/test_pseudo_ops_extension.py b/tests/test_pseudo_ops_extension.py index 9d6c183c4..5836d47e0 100644 --- a/tests/test_pseudo_ops_extension.py +++ b/tests/test_pseudo_ops_extension.py @@ -553,8 +553,10 @@ def test_categorize_as_sdpa_bwd(self): def test_fused_attn_fwd_still_sdpa_fwd(self): """FusedAttnFunc (forward) must still be SDPA_fwd.""" - from TraceLens.PerfModel.op_categories import register_perf_model_categories - from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY + from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, + register_perf_model_categories, + ) registry = dict(OP_CATEGORY_REGISTRY) register_perf_model_categories( @@ -642,8 +644,10 @@ def test_categorization_normalization(self): def test_categorize_as_norm_fwd_bwd(self): """Core categorizer must return NORM_fwd and NORM_bwd.""" - from TraceLens.PerfModel.op_categories import register_perf_model_categories - from TraceLens.PerfModel.torch_op_mapping import OP_CATEGORY_REGISTRY + from TraceLens.PerfModel.torch_op_mapping import ( + OP_CATEGORY_REGISTRY, + register_perf_model_categories, + ) registry = dict(OP_CATEGORY_REGISTRY) register_perf_model_categories( diff --git a/tests/test_torch_op_categorization_registry.py b/tests/test_torch_op_categorization_registry.py index 928089ccb..c9abfece3 100644 --- a/tests/test_torch_op_categorization_registry.py +++ b/tests/test_torch_op_categorization_registry.py @@ -10,14 +10,12 @@ import pytest -from TraceLens.PerfModel.op_categories import ( - build_sheet_category_to_op_names, - register_op_categories, -) from TraceLens.PerfModel.torch_op_mapping import ( OP_CATEGORY_REGISTRY, + build_sheet_category_to_op_names, categorize_torch_op, op_to_perf_model_class_map, + register_op_categories, ) @@ -87,7 +85,9 @@ def test_kernel_name_fallback(kernel_name, expected): def test_registry_covers_every_perf_model_op(): - missing = [name for name in op_to_perf_model_class_map if name not in OP_CATEGORY_REGISTRY] + missing = [ + name for name in op_to_perf_model_class_map if name not in OP_CATEGORY_REGISTRY + ] assert missing == [] From 5045c9cce1bedfab805984219b69ffadde19d4b5 Mon Sep 17 00:00:00 2001 From: Jassani Date: Fri, 8 May 2026 16:28:56 -0400 Subject: [PATCH 07/10] TreePerf: defer metrics source column Keep the categorization cleanup focused by leaving metrics provenance out of this PR. Co-authored-by: Cursor --- TraceLens/TreePerf/tree_perf.py | 7 ------- 1 file changed, 7 deletions(-) diff --git a/TraceLens/TreePerf/tree_perf.py b/TraceLens/TreePerf/tree_perf.py index a838d845c..b9bf3f15f 100644 --- a/TraceLens/TreePerf/tree_perf.py +++ b/TraceLens/TreePerf/tree_perf.py @@ -1973,11 +1973,6 @@ def list_to_tuple(obj): "External id": args.get("External id"), "duration_us": event.get("dur"), "has_perf_model": has_own_perf_model or is_sole_bwd, - "metrics_source": ( - "own_perf_model" - if has_own_perf_model - else "linked_forward_bwd" if is_sole_bwd else "kernel_time_only" - ), "overlapping_kernel_names": event.get("overlapping_kernel_names"), "overlapping_kernels_details": event.get("overlapping_kernels_details"), "overlap_pct": event.get("overlap_pct"), @@ -2142,7 +2137,6 @@ def list_to_tuple(obj): ["Input Dims", "Input type", "Input Strides", "Concrete Inputs"] ) col_order.extend(["duration_us", "has_perf_model", "is_recompute"]) - col_order.append("metrics_source") if include_perf_metrics: col_order.extend(perf_cols) col_order.append("perf_params") @@ -2191,7 +2185,6 @@ def summarize_df_unified_perf_table( grouping_cols = [ "name", "op category", - "metrics_source", "process_name", "process_label", "thread_name", From 8390f4c8a4becfc4891739cef7811cf42ec4dbd2 Mon Sep 17 00:00:00 2001 From: Adeem Jassani Date: Thu, 28 May 2026 13:53:01 -0400 Subject: [PATCH 08/10] PerfModel: post-rebase fixups for FusedLnModulateBackward + sanity test FusedLnModulateBackward (added on main in #633) inherits Normalization, which now declares category='NORM_fwd' on the base. The backward sibling needs an explicit category='NORM_bwd' override, matching every other *Bwd Normalization subclass in this file. Adds test_every_perf_model_class_declares_category so any future perf-model class missing a 'category' attribute fails with a clear, named error rather than a cryptic ValueError from registry construction at import time. Co-authored-by: Cursor --- TraceLens/PerfModel/perf_model.py | 2 ++ .../test_torch_op_categorization_registry.py | 23 ++++++++++++++++++- 2 files changed, 24 insertions(+), 1 deletion(-) diff --git a/TraceLens/PerfModel/perf_model.py b/TraceLens/PerfModel/perf_model.py index ce867ee0e..d7a670051 100644 --- a/TraceLens/PerfModel/perf_model.py +++ b/TraceLens/PerfModel/perf_model.py @@ -5630,6 +5630,8 @@ class FusedLnModulateBackward(Normalization): + B·H·bpe (read modulation_grad) """ + category = "NORM_bwd" + @staticmethod def get_param_details(event): args_input_dims = event["args"]["Input Dims"] diff --git a/tests/test_torch_op_categorization_registry.py b/tests/test_torch_op_categorization_registry.py index c9abfece3..e2f3edde7 100644 --- a/tests/test_torch_op_categorization_registry.py +++ b/tests/test_torch_op_categorization_registry.py @@ -1,5 +1,5 @@ ############################################################################### -# Copyright (c) 2024 - 2025 Advanced Micro Devices, Inc. All rights reserved. +# Copyright (c) 2024 - 2026 Advanced Micro Devices, Inc. All rights reserved. # # See LICENSE for license information. ############################################################################### @@ -14,6 +14,7 @@ OP_CATEGORY_REGISTRY, build_sheet_category_to_op_names, categorize_torch_op, + get_perf_model_category, op_to_perf_model_class_map, register_op_categories, ) @@ -91,6 +92,26 @@ def test_registry_covers_every_perf_model_op(): assert missing == [] +def test_every_perf_model_class_declares_category(): + """Every perf model class must declare a non-None ``category`` attribute. + + Guards against new perf-model PRs that forget to set ``category`` (or + inherit it from a base that does), which would otherwise raise + ``ValueError`` from ``build_op_category_registry`` at import time with + no indication of which class caused it. + """ + missing = [] + for op_name, perf_model_class in op_to_perf_model_class_map.items(): + category = getattr(perf_model_class, "category", None) + if category is None: + missing.append((op_name, perf_model_class.__name__)) + assert missing == [], ( + "perf model classes missing a non-None 'category' attribute: " + f"{missing}. Declare 'category' on the class (or on a base class) " + "so OP_CATEGORY_REGISTRY can resolve it." + ) + + def test_sheet_category_to_op_names_keeps_legacy_sheet_view(): sheet_category_to_op_names = build_sheet_category_to_op_names( op_to_perf_model_class_map From 4de81145dcbc63e25fc5f932681b367037726d2a Mon Sep 17 00:00:00 2001 From: Adeem Jassani Date: Thu, 28 May 2026 14:06:19 -0400 Subject: [PATCH 09/10] PerfModel: drop dead concat_and_cache_mla extension class Class had pass-only stubs and was already commented out in the pseudo-op registration map. The op is now categorized via CATEGORY_ONLY_OP_MAPPING. Confirmed safe to remove by @devalshahamd in PR review. Co-authored-by: Cursor --- .../extensions/perf_model_extensions.py | 19 ------------------- .../extensions/pseudo_ops_perf_utils.py | 2 -- 2 files changed, 21 deletions(-) diff --git a/TraceLens/PerfModel/extensions/perf_model_extensions.py b/TraceLens/PerfModel/extensions/perf_model_extensions.py index cfa4894b4..e2b890325 100644 --- a/TraceLens/PerfModel/extensions/perf_model_extensions.py +++ b/TraceLens/PerfModel/extensions/perf_model_extensions.py @@ -454,25 +454,6 @@ def get_compute_precision(self): return torch_dtype_map(dtype) if dtype else None -class concat_and_cache_mla(BinaryElementwise): - """Performance model for _C_cache_ops::concat_and_cache_mla.""" - - def __init__(self, event, arch=None, python_path=None): - self.event = event - self.arch = arch - self.param_details = self.get_param_details(event) - - @staticmethod - def get_param_details(event): - pass - - def flops(self): - pass - - def bytes(self): - pass - - class vllm_triton_per_token_group_quant_fp8(GroupQuant): """ Performance model for vllm::triton_per_token_group_quant_fp8. diff --git a/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py b/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py index b2b944b2b..5ee4e3394 100644 --- a/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py +++ b/TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py @@ -67,8 +67,6 @@ def get_pseudo_op_mappings(): "vllm::rocm_unquantized_gemm": perf_model_extensions.vllm_rocm_unquantized_gemm, "aiter::gemm_a16w16_atomic_": perf_model_extensions.gemm_a16w16_atomic_, "sglang_profiler::batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant_batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant_464": perf_model_extensions.batched_gemm_a8w8, - ## Cache ops - ##"_C_cache_ops::concat_and_cache_mla": perf_model_extensions.concat_and_cache_mla, ## Quantization ops "vllm::triton_per_token_group_quant_fp8": perf_model_extensions.vllm_triton_per_token_group_quant_fp8, ## Activation ops From 552e90c208d6fbf9df0cae0f724b50edb8da18e1 Mon Sep 17 00:00:00 2001 From: Adeem Jassani Date: Thu, 28 May 2026 15:24:20 -0400 Subject: [PATCH 10/10] docs: add custom perf model recipe Expand perf report column docs with guidance for adding native perf models, using extension files for private/customer-specific ops, and registering category-only ops. Co-authored-by: Cursor --- docs/perf_report_columns.md | 109 ++++++++++++++++++++++++++++++------ 1 file changed, 93 insertions(+), 16 deletions(-) diff --git a/docs/perf_report_columns.md b/docs/perf_report_columns.md index 34e2b5efa..f96687a95 100644 --- a/docs/perf_report_columns.md +++ b/docs/perf_report_columns.md @@ -16,8 +16,9 @@ This document provides detailed explanations of the columns in each sheet of the - [2.2 ops_summary_by_category](#22-ops_summary_by_category-most-aggregated) - [2.3 ops_summary](#23-ops_summary-by-operation-name) - [2.4 ops_unique_args](#24-ops_unique_args-most-detailed) -3. [Performance Metrics Sheets (Roofline Analysis)](#3-performance-metrics-sheets-roofline-analysis) +3. [Unified Perf Summary and Roofline Metrics](#3-unified-perf-summary-and-roofline-metrics) - [Understanding the Metrics Pipeline](#understanding-the-metrics-pipeline) + - [Preferred Output: unified_perf_summary](#preferred-output-unified_perf_summary) - [Operation Parameters Reference](#operation-parameters-reference) - [Why This Matters: Roofline Analysis](#why-this-matters-roofline-analysis) 4. [Collective Communication Analysis](#4-collective-communication-analysis) @@ -36,8 +37,9 @@ The performance report Excel file contains multiple sheets analyzing different a 3. **ops_summary_by_category** - Operations summarized by category 4. **ops_summary** - Operations summarized by name 5. **ops_unique_args** - Operations summarized by unique argument combinations +6. **unified_perf_summary** - Preferred sheet for modeled FLOPs/bytes, runtime TFLOPS/s/TB/s, compute spec, and roofline metrics when available -Additional sheets may include op-specific analysis (GEMM, SDPA_fwd, CONV_fwd, etc.), kernel summary, short kernels, and collective analysis. +Additional sheets may include legacy op-specific analysis (GEMM, SDPA_fwd, CONV_fwd, etc.), kernel summary, short kernels, and collective analysis. **Unit Conventions**: - **Time**: All times from the trace are in **microseconds (µs)** unless explicitly stated otherwise (e.g., `time ms` in `gpu_timeline` is in milliseconds) @@ -583,9 +585,9 @@ This is particularly useful for creating standalone reproducers or benchmarking --- -## 3. Performance Metrics Sheets (Roofline Analysis) +## 3. Unified Perf Summary and Roofline Metrics -For certain operation categories (GEMM, CONV, SDPA, UnaryElementwise, BinaryElementwise), TraceLens generates additional sheets with **roofline model metrics**. These sheets help you understand how efficiently operations are using the GPU's computational and memory bandwidth capabilities. +`unified_perf_summary` is the main TraceLens sheet for perf-model and roofline analysis. It helps you understand how efficiently operations are using the GPU's computational and memory bandwidth capabilities. The report still emits older per-category performance sheets for compatibility, but new analysis should generally start from `unified_perf_summary`. **Important Context**: While hardware counter profilers like `rocprof compute` and `nsight compute` reveal what the GPU actually executed—including effects of padding, redundant memory movement, and cache behavior—TraceLens focuses on the useful work dictated by operator semantics. Used together, these two perspectives provide a richer picture: hardware counters expose low-level execution characteristics, while TraceLens reveals the efficiency of the computation in context. @@ -741,15 +743,39 @@ addmm(6144×2048 × 2048×8192) 762 0.59 ~203 Compute-bound, 58 **Important**: Arithmetic intensity (FLOPs/Byte) determines whether an operation is compute-bound or memory-bound. The percentage of peak achieved indicates optimization quality within that constraint. See [NVIDIA's GEMM Performance Guide](https://docs.nvidia.com/deeplearning/performance/dl-performance-matrix-multiplication/index.html) for more details on this distinction. -### What These Sheets Contain +### Preferred Output: `unified_perf_summary` -Performance metrics sheets are generated for operations with available performance models: +`unified_perf_summary` is the preferred sheet for perf-model output. It combines modeled ops and leaf GPU-launching ops into one table, grouped by operation name and input arguments. Start here before using the older per-category sheets. When `--detect_recompute` is enabled, the table also includes `is_recompute` so recomputed and non-recomputed instances of the same op can be split. + +Rows with a registered performance model include: +- **Static metrics**: GFLOPS, Data Moved (MB), FLOPS/Byte (calculated once from parameters) +- **Runtime metrics**: Kernel Time (µs), TFLOPS/s, TB/s (statistics across all occurrences) +- **Compute Spec**: Combined compute type and precision (e.g., `matrix_bf16`, `vector_fp32`) +- **Roofline metrics** (when `--gpu_arch_json_path` provided): Roofline Time (µs), Pct Roofline - compares achieved time to theoretical roofline bound + +`Compute Spec` is produced by the perf model and combines: +- Compute type: `matrix_*` for matrix-core style work such as GEMM, CONV, and SDPA; `vector_*` for vector/SIMD-style work such as elementwise ops. +- Precision: `fp8`, `fp16`, `bf16`, `fp32`, `fp64`, or another dtype supported by the model. + +When `--gpu_arch_json_path` is provided, TraceLens uses the GPU architecture file to add roofline columns. The architecture file provides `mem_bw_gbps` plus `max_achievable_tflops` values keyed by compute spec (for example `matrix_bf16` or `vector_fp32`). See [GPU Architecture Specifications](../examples/gpu_arch_example.md) for the expected format. + +Rows without a registered performance model can still appear in `unified_perf_summary` if they are leaf ops that launch GPU kernels. These rows have timing and categorization information, but FLOPs/bytes-derived columns are unavailable because TraceLens does not know the theoretical work for the op. + +`op_category_extension` only controls the `op category` label for category-only ops. To populate GFLOPS, Data Moved, FLOPS/Byte, TFLOPS/s, and TB/s, add a native perf model or provide one through `perf_model_extension`. + +**Code Reference**: Generated by `TreePerfAnalyzer.build_df_unified_perf_table()` and `TreePerfAnalyzer.summarize_df_unified_perf_table()`. + +### Legacy Per-Category Performance Sheets + +The report also emits older per-category performance sheets for compatibility and focused analysis. These sheets are useful when a workflow already expects separate tabs such as `GEMM` or `SDPA_fwd`, but new analysis should generally start from `unified_perf_summary`. + +Legacy per-category sheets are generated for operations with available performance models: - **GEMM**: Matrix multiply operations (addmm, mm, bmm, baddbmm, etc.). See [GEMMs in AI Workloads](./conceptual/aimodels_gemms.md) for how model dimensions (batch size, sequence length, hidden dimension) map to GEMM shapes. - **CONV_fwd / CONV_bwd**: Convolution operations - **SDPA_fwd / SDPA_bwd**: Scaled dot-product attention - **UnaryElementwise / BinaryElementwise**: Element-wise operations -Each sheet contains: +Each legacy sheet contains: - All columns from `ops_unique_args` (operation name, arguments, occurrences, etc.) - **Static metrics**: GFLOPS, Data Moved (MB), FLOPS/Byte (calculated once from parameters) - **Compute Spec**: Combined compute type and precision (e.g., `matrix_bf16`, `vector_fp32`). This indicates: @@ -762,7 +788,7 @@ Each sheet contains: - `std_dev`: Variability - high std_dev (>10% of mean) suggests inconsistent performance - `min`, `max`: Range - large spread may indicate outliers worth investigating -**Note**: These sheets contain all the columns from `ops_unique_args`, so you can replay operations from these sheets using the same approach described in [Replaying from Perf Report](#replaying-from-perf-report-no-trace-required). Simply read the desired performance metrics sheet (e.g., `GEMM`, `SDPA_fwd`) instead of `ops_unique_args`. +**Note**: These legacy sheets contain all the columns from `ops_unique_args`, so you can replay operations from these sheets using the same approach described in [Replaying from Perf Report](#replaying-from-perf-report-no-trace-required). Simply read the desired performance metrics sheet (e.g., `GEMM`, `SDPA_fwd`) instead of `ops_unique_args`. **Code Reference**: Generated by `TreePerfAnalyzer.build_df_perf_metrics()` and `TreePerfAnalyzer.summarize_df_perf_metrics()` @@ -841,14 +867,67 @@ For binary ops (add, mul, etc.), two shapes are extracted and broadcasting is ha **FLOPs and Bytes Calculation**: For the specific formulas used to calculate FLOPs and memory traffic from these parameters, see the performance model implementations in `TraceLens/PerfModel/perf_model.py`. -**Extending with Custom Operations**: Performance metrics sheets are only generated for operations that have a registered performance model. If your workload includes custom operations (e.g., from Megatron, vLLM, or other libraries), you can extend TraceLens by: -1. Creating a performance model class for your operation (inherit from `GEMM`, `CONV`, `SDPA`, etc.) -2. Implementing `get_param_details()`, `flops()`, and `bytes()` methods -3. Passing an extension file via `--extension_file` when generating the report +### Extending TraceLens With Custom Perf Models -See `examples/megatron_extension.py` for a complete example of extending TraceLens with custom Megatron operations. +Perf-model output is primarily consumed through `unified_perf_summary`. If an operation only needs to be labeled in that table, use `op_category_extension` instead of adding a perf model. Add a perf model when TraceLens should calculate theoretical FLOPs, bytes moved, arithmetic intensity, TFLOPS/s, or TB/s for the op. -**Note**: If your custom operation is frequently used across multiple projects, we can work to add it as a native operation in TraceLens. Please open an issue or reach out to discuss integration. +For generally useful operators, add the model natively in TraceLens: + +1. Create a performance model class in `TraceLens/PerfModel/perf_model.py` or a focused module under `TraceLens/PerfModel/extensions/`. +2. Inherit from the closest existing base (`GEMM`, `CONV`, `SDPA`, `Normalization`, `UnaryElementwise`, etc.) when possible. +3. Set `category` on the class, for example `category = "GEMM"` or `category = "NORM_bwd"`. Use `bwd_category` only when a forward model intentionally computes linked backward metrics. Use `sheet_category` only when the legacy sheet name should differ from the runtime category. +4. Implement `get_param_details(event)` to parse the trace event's `args` into the parameters needed by the model. +5. Implement or inherit `flops()` and `bytes()`. +6. Make sure the model reports the right compute spec. Existing bases usually provide this; custom bases should implement `get_maf_type()` (for example `matrix` or `vector`) and `get_compute_precision()` so `--gpu_arch_json_path` can look up the matching MAF entry such as `matrix_bf16`. +7. Register the op name in `TraceLens/PerfModel/torch_op_mapping.py` via `op_to_perf_model_class_map`. +8. Add tests for categorization and model behavior. At minimum, make sure `tests/test_torch_op_categorization_registry.py` covers the new op category. + +For private or customer-specific operators that should not live in the shared source tree, use an extension file and pass it with `--extension_file`: + +```python +from TraceLens.PerfModel.perf_model import GEMM, name2bpe + + +class my_custom_gemm(GEMM): + category = "GEMM" + + @staticmethod + def get_param_details(event): + input_dims = event["args"]["Input Dims"] + input_types = event["args"]["Input type"] + M, K = input_dims[0] + _, N = input_dims[1] + return { + "M": M, + "K": K, + "N": N, + "bias": False, + "dtype_A_B": (input_types[0], input_types[1]), + } + + def bytes(self): + dtype_a, dtype_b = self.param_details["dtype_A_B"] + bpe_a = name2bpe(dtype_a) + bpe_b = name2bpe(dtype_b) + # Assume the output uses the same dtype as A for this example. + return super().bytes( + bpe_mat1=bpe_a, + bpe_mat2=bpe_b, + bpe_bias=bpe_a, + bpe_output=bpe_a, + ) + + +perf_model_extension = { + "my_namespace::custom_gemm": my_custom_gemm, +} + +op_category_extension = { + "my_namespace::custom_attention_backward_without_model": "SDPA_bwd", +} +``` + +See `examples/example_megatron_extension.py` for a complete extension example. --- @@ -943,8 +1022,6 @@ Depending on the command-line arguments used when generating the report, additio - **short_kernel_histogram**, **short_kernels_summary**: Analysis of very short kernels (enabled with `--short_kernel_study`) -- **unified_perf_summary**: Unified perf metrics for all ops with perf models or leaf ops that launch GPU kernels. Includes GFLOPS, TFLOPS/s, Data Moved, FLOPS/Byte, TB/s metrics aggregated by unique args. When `--detect_recompute` is enabled, an `is_recompute` column is added to split rows by recompute status. - --- ## 6. Common Analysis Workflows