diff --git a/TraceLens/Reporting/generate_perf_report_pytorch.py b/TraceLens/Reporting/generate_perf_report_pytorch.py index e1ed51a4b..70b2f9630 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch.py @@ -318,6 +318,15 @@ def apply_extension(perf_analyzer, extension_path): op_category_extension, OP_CATEGORY_REGISTRY, ) + if hasattr(extension, "categorize_extension"): + custom_categorizer = getattr(extension, "categorize_extension") + base_categorizer = perf_analyzer.op_categorizer + + def op_categorizer(row): + category = custom_categorizer(row, perf_analyzer) + return category if category is not None else base_categorizer(row) + + perf_analyzer.op_categorizer = op_categorizer if hasattr(extension, "dict_cat2names_extension"): warnings.warn( "dict_cat2names_extension is deprecated and ignored. Use " diff --git a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py index 905193a34..1bc632c13 100644 --- a/TraceLens/Reporting/generate_perf_report_pytorch_inference.py +++ b/TraceLens/Reporting/generate_perf_report_pytorch_inference.py @@ -430,6 +430,15 @@ def apply_extension(perf_analyzer, extension_path): op_category_extension, OP_CATEGORY_REGISTRY, ) + if hasattr(extension, "categorize_extension"): + custom_categorizer = getattr(extension, "categorize_extension") + base_categorizer = perf_analyzer.op_categorizer + + def op_categorizer(row): + category = custom_categorizer(row, perf_analyzer) + return category if category is not None else base_categorizer(row) + + perf_analyzer.op_categorizer = op_categorizer if hasattr(extension, "dict_cat2names_extension"): warnings.warn( "dict_cat2names_extension is deprecated and ignored. Use " diff --git a/docs/how-to/generate-perf-report-pytorch.md b/docs/how-to/generate-perf-report-pytorch.md index e80f83133..7aebd35eb 100644 --- a/docs/how-to/generate-perf-report-pytorch.md +++ b/docs/how-to/generate-perf-report-pytorch.md @@ -167,6 +167,7 @@ can define any of: | `tree_postprocess_extension` | `Callable` | Called with `perf_analyzer.tree`; update the tree post-construction. | | `perf_model_extension` | `dict` | Map op name → custom perf-model class; overrides or extends built-in models. | | `op_category_extension` | `dict` | Map category-only op names to final categories, so an op appears in unified reports without a perf model. | +| `categorize_extension` | `Callable` | Called with `(row, perf_analyzer)`; return a category or `None` to use the default categorizer. | ```bash TraceLens_generate_perf_report_pytorch \ diff --git a/examples/example_megatron_extension.py b/examples/example_megatron_extension.py index d0afee2d1..909e04c76 100644 --- a/examples/example_megatron_extension.py +++ b/examples/example_megatron_extension.py @@ -422,11 +422,25 @@ def inject_pseudo_op( # we also need to -def categorize_extension(row, plugin): +def categorize_extension(row, _perf_analyzer): """ Categorizer plugin to categorize the kernel launchers. """ - if row["name"] in [ + name = row["name"] + grouped_bwd_prefix = "_GroupedLinearBackward->" + synthetic_suffix = " (Synthetic Op)" + if name.startswith(grouped_bwd_prefix) and name.endswith(synthetic_suffix): + kernel_name = name[len(grouped_bwd_prefix) : -len(synthetic_suffix)] + if is_gemm_kernel({"cat": "kernel", "name": kernel_name}): + return "GroupedGEMM_bwd" + if name.endswith(synthetic_suffix) and name.startswith( + ( + "_OperationFuserAutogradFunctionBackward->ln_tma_bwd_kernel", + "_OperationFuserAutogradFunctionBackward->ln_bwd_finalize_kernel", + ) + ): + return "NORM_bwd" + if name in [ "_Linear_fwd_mm", "_LayerNormLinear_fwd_mm", "_LinearBackward_xgrad_mm", @@ -435,9 +449,9 @@ def categorize_extension(row, plugin): "_LayerNormLinearBackward_wgrad_mm", ]: return "GEMM" - if row["name"] == "FusedAttnFunc": + if name == "FusedAttnFunc": return "SDPA_fwd" - if row["name"] == "FusedAttnFuncBackward": + if name == "FusedAttnFuncBackward": return "SDPA_bwd" return None @@ -718,4 +732,6 @@ def get_param_details(event): op_category_extension = { "FusedAttnFuncBackward": "SDPA_bwd", "GroupedGemmBackward": "GroupedGEMM_bwd", + "_GroupedLinear": "GroupedGEMM_fwd", + "_GroupedLinearBackward": "GroupedGEMM_bwd", } diff --git a/tests/test_pseudo_ops_extension.py b/tests/test_pseudo_ops_extension.py index bc3decab8..de1d04565 100644 --- a/tests/test_pseudo_ops_extension.py +++ b/tests/test_pseudo_ops_extension.py @@ -22,6 +22,7 @@ from TraceLens.TreePerf.tree_perf import TreePerfAnalyzer from example_megatron_extension import ( _link_checkpoint_fwd_bwd, + categorize_extension, op_category_extension, perf_model_extension, te_layer_norm_bwd, @@ -601,6 +602,60 @@ def test_fused_attn_fwd_still_sdpa_fwd(self): assert registry["FusedAttnFunc"] == "SDPA_fwd" +@pytest.mark.parametrize( + "name,expected", + [ + ("_GroupedLinear", "GroupedGEMM_fwd"), + ("_GroupedLinearBackward", "GroupedGEMM_bwd"), + ], +) +def test_megatron_category_only_mappings(name, expected): + assert op_category_extension[name] == expected + + +@pytest.mark.parametrize( + "name,expected", + [ + ( + "_GroupedLinearBackward->nvjet_sm103_qrtst (Synthetic Op)", + "GroupedGEMM_bwd", + ), + ( + "_GroupedLinearBackward->Cijk_Ailk_Bljk (Synthetic Op)", + "GroupedGEMM_bwd", + ), + ( + "_GroupedLinearBackward->RR_GEMM_test (Synthetic Op)", + "GroupedGEMM_bwd", + ), + ( + "_OperationFuserAutogradFunctionBackward->ln_tma_bwd_kernel " + "(Synthetic Op)", + "NORM_bwd", + ), + ( + "_OperationFuserAutogradFunctionBackward->ln_bwd_finalize_kernel " + "(Synthetic Op)", + "NORM_bwd", + ), + ], +) +def test_megatron_synthetic_category_mappings(name, expected): + assert categorize_extension({"name": name}, None) == expected + + +@pytest.mark.parametrize( + "name", + [ + "_OperationFuserAutogradFunctionBackward->ln_tma_fwd_kernel (Synthetic Op)", + "_OperationFuserAutogradFunctionBackward->ln_tma_bwd_kernel", + "_GroupedLinearBackward->quantize_kernel (Synthetic Op)", + ], +) +def test_megatron_synthetic_category_mapping_ignores_unmatched_ops(name): + assert categorize_extension({"name": name}, None) is None + + class TestLayerNormFnPerfModel: """Test LayerNormFn / LayerNormFnBackward perf model and categorization.""" diff --git a/tests/test_reporting_utils.py b/tests/test_reporting_utils.py index c90a127de..80e907198 100644 --- a/tests/test_reporting_utils.py +++ b/tests/test_reporting_utils.py @@ -683,6 +683,30 @@ class DummyGemm: assert "aten::mm" in analyzer.op_to_perf_model_class_map +@pytest.mark.parametrize( + "apply_extension", [apply_extension_pytorch, apply_extension_inference] +) +def test_apply_extension_categorizer_hook(tmp_path, apply_extension): + ext_path = tmp_path / "ext.py" + ext_path.write_text(textwrap.dedent(""" + def categorize_extension(row, plugin): + assert plugin is not None + if row["name"] == "custom::op": + return "custom" + return None + """)) + analyzer = SimpleNamespace( + tree=SimpleNamespace(events=[], label_non_gpu_paths=lambda: None), + op_to_perf_model_class_map={}, + op_categorizer=lambda row: "base", + ) + + apply_extension(analyzer, str(ext_path)) + + assert analyzer.op_categorizer({"name": "custom::op"}) == "custom" + assert analyzer.op_categorizer({"name": "unknown::op"}) == "base" + + # --------------------------------------------------------------------------- # trunc / wrapper helpers # ---------------------------------------------------------------------------