Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
763d806
feat(perfmodel): add varlen aiter FlashAttention + aten::_flash_atten…
gphuang May 19, 2026
f2a2480
test(perfmodel): add unit tests for new aiter varlen + aten flash per…
gphuang May 19, 2026
63214a2
fix(perfmodel): account for num_seqs in aiter varlen FA FLOPs (Copilo…
gphuang May 19, 2026
f81d5ac
test(perfmodel): refresh refs for new SDPA attention coverage
gphuang May 25, 2026
585e2f2
test(perfmodel): refresh compare report refs for SDPA coverage
gphuang May 25, 2026
6ab1888
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang May 26, 2026
b8f2b92
docs(perfmodel): move fmha_v3_varlen_fwd mapping note to extensions
gphuang May 26, 2026
c94523b
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang May 28, 2026
80e2551
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang May 29, 2026
fce38b1
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang Jun 4, 2026
21f69d2
test(perfmodel): refresh torch_compile_triton refs for SDPA coverage
gphuang Jun 4, 2026
ad64a55
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang Jun 12, 2026
e5c9790
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang Jul 7, 2026
85d1d03
docs(readthedocs): add minimal Sphinx docs build config
gphuang Jul 7, 2026
92aa02a
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang Jul 9, 2026
4df85aa
fix(perfmodel): restore varlen extension priority and origami dtype f…
gphuang Jul 9, 2026
8dc07cf
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang Jul 10, 2026
c0c5580
Merge branch 'main' into feat/perfmodel/wan22-attention-coverage
gphuang Jul 14, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -580,27 +580,79 @@ class mha_varlen_fwd(InferenceAttention):

class aiter_fmha_v3_varlen_fwd(InferenceAttention):
"""
Performance model for ``aiter::fmha_v3_varlen_fwd`` (inference: sglang / vLLM).
Annotation-aware perf model for ``aiter::fmha_v3_varlen_fwd`` (inference:
sglang / vLLM).

Uses the same chunk statistics as :class:`InferenceAttention` (``annotation``
on the event). Sets ``d_h_v`` from the **v** tensor (``Input Dims[2]``) so MLA
shapes with differing Q/K vs V head dims are modeled correctly.

Unparseable annotation yields :meth:`InferenceAttention.no_perf_param_details`
(see base class); no packed-tensor fallback.
If annotation parsing fails, this class falls back to the core shape-based
SDPA model so regular training traces still get FLOPs/bytes coverage while
annotated inference traces use the more accurate per-request sequence stats.
"""

category = "SDPA_fwd"
bwd_category = None

def __init__(self, event, arch=None, python_path=None, enable_origami=False):
self.enable_origami = enable_origami
super().__init__(event, arch, python_path)

@staticmethod
def _core_param_details(event):
from TraceLens.PerfModel import perf_model

params = perf_model.aiter__fmha_v3_varlen_fwd.get_param_details(event).copy()
params["_fallback_core"] = True
return params

def _core_model(self):
from TraceLens.PerfModel import perf_model

return perf_model.aiter__fmha_v3_varlen_fwd(
self.event,
self.arch,
self.python_path,
enable_origami=self.enable_origami,
)

@staticmethod
def get_param_details(event):
params = InferenceAttention.get_param_details(event)
if params.get("_no_perf"):
return params
try:
return aiter_fmha_v3_varlen_fwd._core_param_details(event)
except (ValueError, IndexError, KeyError, TypeError):
return params
args = event.get("args") or {}
dims = args.get("Input Dims") or []
if len(dims) > 2 and len(dims[2]) >= 1:
params["d_h_v"] = dims[2][-1]
return params

def flops(self):
if self.param_details.get("_fallback_core"):
return self._core_model().flops()
return super().flops()

def bytes(self, bytes_per_element=None):
if self.param_details.get("_fallback_core"):
if bytes_per_element is None:
return self._core_model().bytes()
return self._core_model().bytes(bytes_per_element)
return super().bytes(bytes_per_element)

def get_compute_precision(self):
if self.param_details.get("_fallback_core"):
return self._core_model().get_compute_precision()
return super().get_compute_precision()

def get_simulation_time(self):
if self.param_details.get("_fallback_core"):
return self._core_model().get_simulation_time()
return None


class aiter_paged_attention_ragged(InferenceAttention):
"""
Expand Down
2 changes: 2 additions & 0 deletions TraceLens/PerfModel/extensions/pseudo_ops_perf_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,8 @@ def get_pseudo_op_mappings():
# Attention pseudo ops
"vllm::unified_attention_with_output": attention_perf_model_extensions.vllm_unified_attention_with_output,
"aiter::mha_varlen_fwd": attention_perf_model_extensions.mha_varlen_fwd,
# Prefer the annotation-aware inference extension when annotations are
# present; it falls back to the core SDPA model for plain training rows.
"aiter::fmha_v3_varlen_fwd": attention_perf_model_extensions.aiter_fmha_v3_varlen_fwd,
"aiter::mha_batch_prefill": attention_perf_model_extensions.aiter_mha_batch_prefill,
"sglang_profiler::attention_paged_attention_ragged": attention_perf_model_extensions.aiter_paged_attention_ragged,
Expand Down
Loading
Loading