[rocprofiler-compute] Type annotation patch for analysis_db.py (#981)

This commit is contained in:
xuchen-amd
2025-09-23 17:05:37 -04:00
committed by GitHub
parent 872f0aed0c
commit c3054c00b1
18 changed files with 44 additions and 32 deletions
@@ -24,7 +24,6 @@ extend-exclude = [
# deprecated files
"src/rocprof_compute_profile/profiler_rocprof_v1.py",
"src/rocprof_compute_profile/profiler_rocprof_v2.py",
"src/rocprof_compute_analyze/analysis_db.py",
"src/utils/db_connector.py",
# WIP files
"src/rocprof_compute_tui/widgets/splitter.py"
@@ -23,7 +23,6 @@
##############################################################################
from pathlib import Path
# NB: Creating a new module to share global vars across modules
@@ -23,7 +23,6 @@
##############################################################################
from rocprof_compute_analyze.analysis_base import OmniAnalyze_Base
from utils import file_io, parser, tty
from utils.kernel_name_shortener import kernel_name_shortener
@@ -26,6 +26,7 @@ import ast
import json
import re
from pathlib import Path
from typing import Any, Callable, Optional, Union
import astunparse
import pandas as pd
@@ -37,9 +38,9 @@ from utils import rocpd_data
from utils.analysis_orm import Database, get_views
from utils.logger import console_debug, console_error, console_warning, demarcate
from utils.parser import (
BUILD_IN_VARS,
PC_SAMPLING_NOT_ISSUE_PREFIX,
CodeTransformer,
BUILD_IN_VARS,
to_avg,
to_concat,
to_int,
@@ -66,7 +67,7 @@ class db_analysis(OmniAnalyze_Base):
# Required child methods
# -----------------------
@demarcate
def pre_processing(self):
def pre_processing(self) -> None:
"""Perform any pre-processing steps prior to analysis."""
super().pre_processing()
if self._profiling_config.get("format_rocprof_output") != "rocpd":
@@ -101,7 +102,7 @@ class db_analysis(OmniAnalyze_Base):
self._roofline_data_per_workload = self.calc_roofline_data()
@demarcate
def run_analysis(self):
def run_analysis(self) -> None:
"""Run CLI analysis."""
super().run_analysis()
@@ -210,8 +211,8 @@ class db_analysis(OmniAnalyze_Base):
console_debug("Completed writing database")
console_warning(f"Created file: {db_name}")
def calc_roofline_ceilings(self):
roofline_ceilings_per_workload = dict()
def calc_roofline_ceilings(self) -> dict[str, dict[str, Any]]:
roofline_ceilings_per_workload: dict[str, dict[str, Any]] = {}
for workload_path in self._runs.keys():
if not (Path(workload_path) / "roofline.csv").exists():
@@ -221,7 +222,7 @@ class db_analysis(OmniAnalyze_Base):
roofline_dict = (
pd.read_csv(f"{workload_path}/roofline.csv").iloc[0].to_dict()
)
keys = list()
keys: list[str] = []
for mem_level in CACHE_HIERARCHY:
keys.append(f"{mem_level}Bw")
for dtype in SUPPORTED_DATATYPES[
@@ -247,8 +248,8 @@ class db_analysis(OmniAnalyze_Base):
console_debug("Collected roofline ceilings")
return roofline_ceilings_per_workload
def calc_pc_sampling_data(self):
pc_sampling_data_per_workload = dict()
def calc_pc_sampling_data(self) -> dict[str, pd.DataFrame]:
pc_sampling_data_per_workload: dict[str, pd.DataFrame] = {}
for workload_path in self._runs.keys():
if not (Path(workload_path) / "ps_file_results.json").exists():
@@ -289,16 +290,18 @@ class db_analysis(OmniAnalyze_Base):
for pc_sample in pc_sampling_stochastic + pc_sampling_host_trap
])
def custom_aggregator(column_name):
def custom_aggregator(
column_name: str,
) -> Callable[[pd.Series], Union[int, dict[str, int], None]]:
if column_name == "count_issued":
def aggregator(series):
def aggregator(series: pd.Series) -> Optional[int]:
return None if series.isnull().all() else series.sum()
return aggregator
if column_name == "count_stalled":
def aggregator(series):
def aggregator(series: pd.Series) -> Optional[int]:
if series.isnull().all():
return None
return series.count() - series.sum()
@@ -306,7 +309,7 @@ class db_analysis(OmniAnalyze_Base):
return aggregator
if column_name == "stall_reason":
def aggregator(series):
def aggregator(series: pd.Series) -> Optional[dict[str, int]]:
if series.isnull().all():
return None
cleaned_series = series.dropna().str[
@@ -352,7 +355,13 @@ class db_analysis(OmniAnalyze_Base):
return pc_sampling_data_per_workload
@staticmethod
def evaluate(name, value, pmc_df, sys_info, parse=False):
def evaluate(
name: str,
value: str,
pmc_df: pd.DataFrame,
sys_info: dict[str, Any], # noqa ANN401
parse: bool = False,
) -> Any: # noqa ANN401
if parse:
value = re.sub(
r"\$([0-9A-Za-z_]+)",
@@ -397,7 +406,7 @@ class db_analysis(OmniAnalyze_Base):
console_warning(f"Failed to evaluate expression for {name}: {value} - {e}")
return None
def calc_expressions(self):
def calc_expressions(self) -> dict[str, pd.DataFrame]:
values_data_per_workload = self._values_data_per_workload.copy()
for workload_path in self._runs.keys():
@@ -438,9 +447,11 @@ class db_analysis(OmniAnalyze_Base):
console_debug("Calculated metric values")
return values_data_per_workload
def calc_metrics_data(self):
metrics_info_data_per_workload = dict()
values_data_per_workload = dict()
def calc_metrics_data(
self,
) -> tuple[dict[str, pd.DataFrame], dict[str, pd.DataFrame]]:
metrics_info_data_per_workload: dict[str, pd.DataFrame] = {}
values_data_per_workload: dict[str, pd.DataFrame] = {}
for workload_path in self._runs.keys():
gfx_arch = self._runs[workload_path].sys_info.iloc[0]["gpu_arch"]
@@ -504,8 +515,8 @@ class db_analysis(OmniAnalyze_Base):
console_debug("Collected metrics data")
return metrics_info_data_per_workload, values_data_per_workload
def calc_dispatch_data(self):
dispatch_data_per_workload = dict()
def calc_dispatch_data(self) -> dict[str, pd.DataFrame]:
dispatch_data_per_workload: dict[str, pd.DataFrame] = {}
for workload_path in self._runs.keys():
dispatch_df = pd.DataFrame([
@@ -522,7 +533,7 @@ class db_analysis(OmniAnalyze_Base):
console_debug("Calculated dispatch data")
return dispatch_data_per_workload
def apply_pmc_filters(self):
def apply_pmc_filters(self) -> dict[str, pd.DataFrame]:
pmc_df_per_workload = self._pmc_df_per_workload.copy()
for workload_path, pmc_df in pmc_df_per_workload.items():
@@ -559,8 +570,8 @@ class db_analysis(OmniAnalyze_Base):
console_debug("Applied analysis mode filters")
return pmc_df_per_workload
def calc_roofline_data(self):
roofline_data_per_workload = dict()
def calc_roofline_data(self) -> dict[str, pd.DataFrame]:
roofline_data_per_workload: dict[str, pd.DataFrame] = {}
for workload_path in self._runs.keys():
pmc_df = self._pmc_df_per_workload[workload_path].copy()
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import argparse
from typing import Any, Optional
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import argparse
from typing import Any, Optional
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import argparse
from typing import Any, Optional
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import argparse
from typing import Any, Optional
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import argparse
from typing import Any, Optional
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import argparse
from typing import Any, Optional
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
from typing import Any, Optional, Union
import pandas as pd
@@ -23,7 +23,6 @@
##############################################################################
from textual.app import ComposeResult
from textual.containers import Container, Horizontal
from textual.screen import ModalScreen
@@ -21,6 +21,7 @@
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
##############################################################################el
from typing import Any, Optional
from sqlalchemy import (
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import re
from collections import OrderedDict
from pathlib import Path
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
from typing import Any
from dash import html
@@ -22,6 +22,7 @@
# THE SOFTWARE.
##############################################################################
import argparse
import ast
import json
@@ -23,7 +23,6 @@
##############################################################################
import csv
from dataclasses import dataclass
from pathlib import Path
@@ -23,11 +23,6 @@
##############################################################################
#
# Define all common data storage classes,
# predifned dict and global functions.
#
from collections import OrderedDict
from dataclasses import dataclass, field
from typing import Any