Skip to content
Open
2 changes: 0 additions & 2 deletions TraceLens/PerfModel/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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("_")]
3 changes: 1 addition & 2 deletions TraceLens/PerfModel/extensions/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -92,5 +92,4 @@
"custom_ar_qr_all_reduce",
# Utility functions
"get_pseudo_op_mappings",
"get_pseudo_op_categories",
]
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,9 @@ class InferenceAttention:
``params.get("_no_perf")`` is true.
"""

category = "InferenceAttention"
bwd_category = None

REQUIRED_PARAM_KEYS = (
"B",
"N_Q",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,9 @@


class CustomCollective:
category = "CustomCollective"
bwd_category = None

pass


Expand Down
9 changes: 9 additions & 0 deletions TraceLens/PerfModel/extensions/moe_perf_model_extensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
"""
Expand Down Expand Up @@ -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):
"""
Expand Down
22 changes: 3 additions & 19 deletions TraceLens/PerfModel/extensions/perf_model_extensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -451,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.
Expand Down
30 changes: 0 additions & 30 deletions TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -103,31 +101,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
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,15 @@
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
sheet_category = "RMSNorm"


class aiter_rms_norm(RMSNorm):
Expand Down
Loading