[Refactor][Workloads] Align workloads/ops/ naming with tileops/ops/ layout (#939)

## Summary

Align `workloads/ops/` file naming and directory structure 1:1 with the
post-#928 `tileops/ops/` layout. Pure file-move/rename refactor — no
workload logic changes.

- Drop `_fwd` suffix from mamba, attention, deltanet, and gated_deltanet
workloads
- Rename `mean_pooling_ops.py` to `mean_pooling.py`
- Move attention workloads into `workloads/ops/attention/` subpackage
- Update all imports across the repo

Closes #931

## Test plan

- [x] **AC-1**: All `from workloads.ops.<name>` imports updated for
moved/renamed files — verified by importing all 20 renamed/moved
workload classes and grep confirming zero old import paths remain
- [x] **AC-2**: `python -c "from workloads.ops.attention import ..."`
resolves for all 15 moved attention workload classes
- [x] **AC-3**: Full test suite passes — 2361 passed, 22 skipped, 0
failed (233.76s at commit 4dc54b7)
- [x] **AC-4**: No orphaned files remain after migration — all 11 old
filenames confirmed absent; `mhc_post.py`, `mhc_pre.py` at
`workloads/ops/` root; `nsa_utils.py` at `workloads/` root

## Follow-up

No follow-up issues or suggestions.

---------

Co-authored-by: Ibuki 🍃 — a wind born from Claude Opus <Ibuki-wind@users.noreply.github.com>
This commit is contained in:
Cao Ying 2026-04-13 10:44:06 +08:00 committed by GitHub
parent b77f40c762
commit 6e0b507b42
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
160 changed files with 629 additions and 664 deletions

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import DeepSeekSparseAttentionDecodeWithKVCacheFwdOp
from workloads.ops.deepseek_dsa_decode import DsaDecodeTest
from workloads.attention.deepseek_dsa_decode import DsaDecodeTest
class _DsaDecodeTestBaseline(DsaDecodeTest):

View File

@ -7,7 +7,7 @@ from einops import einsum, rearrange
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import MultiHeadLatentAttentionDecodeWithKVCacheFwdOp
from workloads.ops.deepseek_mla_decode import MlaDecodeTest
from workloads.attention.deepseek_mla_decode import MlaDecodeTest
class _MlaDecodeTestBaseline(MlaDecodeTest):

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import NSAFwdVarlenOp
from workloads.ops.deepseek_nsa_fwd import NsaFwdTest
from workloads.attention.deepseek_nsa import NsaFwdTest
class _NsaFwdTestBaseline(NsaFwdTest):

View File

@ -5,8 +5,8 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import NSACmpFwdVarlenOp
from workloads.attention.deepseek_nsa_cmp import NsaCmpFwdTest
from workloads.nsa_utils import prepare_chunk_offsets
from workloads.ops.deepseek_nsa_cmp_fwd import NsaCmpFwdTest
def _parallel_nsa_compression_fwd_pytorch(test, q, k_cmp, v_cmp, block_size, scale, offsets):

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import NSATopkVarlenOp
from workloads.ops.deepseek_nsa_topk import NsaTopkTest
from workloads.attention.deepseek_nsa_topk import NsaTopkTest
def _nsa_topk_torch(test, q, k_cmp, lse, block_counts, block_size, scale,

View File

@ -6,7 +6,7 @@ from torch.nn import functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GroupedQueryAttentionBwdOp, GroupedQueryAttentionFwdOp
from workloads.ops.gqa import (
from workloads.attention.gqa import (
GqaBwdTest,
GqaFwdTest,
)

View File

@ -7,7 +7,7 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GroupedQueryAttentionDecodeWithKVCacheFwdOp
from workloads.ops.gqa_decode import GqaDecodeTest
from workloads.attention.gqa_decode import GqaDecodeTest
class _GqaDecodeTestBaseline(GqaDecodeTest):

View File

@ -8,7 +8,7 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GroupedQueryAttentionDecodePagedWithKVCacheFwdOp
from workloads.ops.gqa_decode_paged import GqaDecodePagedTest
from workloads.attention.gqa_decode_paged import GqaDecodePagedTest
class _GqaDecodePagedTestBaseline(GqaDecodePagedTest):

View File

@ -7,7 +7,7 @@ from torch.nn import functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GqaSlidingWindowFwdOp
from workloads.ops.gqa_sliding_window_fwd import GqaSlidingWindowFwdTest
from workloads.attention.gqa_sliding_window import GqaSlidingWindowFwdTest
class GqaSlidingWindowFwdBenchmark(BenchmarkBase):

View File

@ -7,7 +7,7 @@ from torch.nn import functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GqaSlidingWindowVarlenFwdOp
from workloads.ops.gqa_sliding_window_varlen_fwd import GqaSlidingWindowVarlenFwdTest
from workloads.attention.gqa_sliding_window_varlen import GqaSlidingWindowVarlenFwdTest
_GQA_SLIDING_WINDOW_VARLEN_FWD_BENCH_PARAMS = [
pytest.param(1, [3000], [3000], 32, 8, 128, True, -1, -1, torch.float16, False, id="single-seq-causal"),

View File

@ -5,8 +5,8 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import MeanPoolingForwardOp
from workloads.attention.mean_pooling import MeanPoolingTest
from workloads.nsa_utils import prepare_chunk_indices
from workloads.ops.mean_pooling_ops import MeanPoolingTest
class _MeanPoolingTestBaseline(MeanPoolingTest):

View File

@ -6,7 +6,7 @@ from torch.nn import functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import MultiHeadAttentionBwdOp, MultiHeadAttentionFwdOp
from workloads.ops.mha import (
from workloads.attention.mha import (
MhaBwdTest,
MhaFwdTest,
)

View File

@ -7,7 +7,7 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import MultiHeadAttentionDecodeWithKVCacheFwdOp
from workloads.ops.mha_decode import MhaDecodeTest
from workloads.attention.mha_decode import MhaDecodeTest
class _MhaDecodeTestBaseline(MhaDecodeTest):

View File

@ -7,7 +7,7 @@ import torch.nn.functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import MultiHeadAttentionDecodePagedWithKVCacheFwdOp
from workloads.ops.mha_decode_paged import MhaDecodePagedTest
from workloads.attention.mha_decode_paged import MhaDecodePagedTest
class _MhaDecodePagedTestBaseline(MhaDecodePagedTest):

View File

@ -26,8 +26,8 @@ from tileops.kernels.elementwise import (
_make_unary_explicit,
)
from tileops.ops.elementwise import ErfOp, GeluOp, MishOp, ReluOp
from workloads.activation import ReluTest
from workloads.base import FixtureBase
from workloads.ops.activation import ReluTest
# ---------------------------------------------------------------------------
# LLM-realistic shapes (LLaMA-family defaults)

View File

@ -8,8 +8,8 @@ from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.ada_layer_norm import AdaLayerNormFwdOp
from tileops.ops.norm.ada_layer_norm_zero import AdaLayerNormZeroFwdOp
from workloads.ops.ada_layer_norm import AdaLayerNormTest
from workloads.ops.ada_layer_norm_zero import AdaLayerNormZeroTest
from workloads.ada_layer_norm import AdaLayerNormTest
from workloads.ada_layer_norm_zero import AdaLayerNormZeroTest
_ADA_OP_NAME = "AdaLayerNormFwdOp"
_ADA_ZERO_OP_NAME = "AdaLayerNormZeroFwdOp"

View File

@ -10,7 +10,7 @@ import torch
from benchmarks.benchmark import BenchmarkReport, ManifestBenchmark, workloads_to_params
from tileops.ops.reduction.argmax import ArgmaxFwdOp
from tileops.ops.reduction.argmin import ArgminFwdOp
from workloads.ops.argreduce import ArgmaxTest, ArgminTest
from workloads.argreduce import ArgmaxTest, ArgminTest
_ARGMAX_OP = "ArgmaxFwdOp"
_ARGMIN_OP = "ArgminFwdOp"

View File

@ -15,7 +15,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.batch_norm import BatchNormBwdOp, BatchNormFwdOp
from workloads.ops.batch_norm import BatchNormBwdTest, BatchNormFwdTest
from workloads.batch_norm import BatchNormBwdTest, BatchNormFwdTest
_FWD_OP_NAME = "BatchNormFwdOp"
_BWD_OP_NAME = "BatchNormBwdOp"

View File

@ -20,7 +20,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops.elementwise import AddOp, WhereOp
from workloads.base import FixtureBase
from workloads.ops.binary_arith import AddSameShapeTest
from workloads.binary_arith import AddSameShapeTest
# ---------------------------------------------------------------------------
# LLM-realistic shapes (LLaMA-family defaults)

View File

@ -21,7 +21,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import DeltaNetBwdOp, DeltaNetFwdOp, DeltaNetOp
from workloads.base import FixtureBase
from workloads.ops.deltanet_fwd import DeltaNetFwdTest
from workloads.deltanet import DeltaNetFwdTest
def _differentiable_fwd(q, k, v, beta, chunk_size):

View File

@ -6,7 +6,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import DeltaNetDecodeOp
from workloads.base import FixtureBase
from workloads.ops.deltanet_recurrence import DeltaNetDecodeTest
from workloads.deltanet import DeltaNetDecodeTest
def deltanet_decode_torch(

View File

@ -7,9 +7,12 @@ import torch.nn.functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops.engram import EngramGateConvBwdOp, EngramGateConvFwdOp
from tileops.ops.engram_decode import EngramDecodeOp
from workloads.ops.engram_bwd import EngramGateConvBwdTest
from workloads.ops.engram_decode import EngramDecodeTest
from workloads.ops.engram_fwd import CONV_KERNEL_SIZE, EngramGateConvFwdTest
from workloads.engram import (
CONV_KERNEL_SIZE,
EngramDecodeTest,
EngramGateConvBwdTest,
EngramGateConvFwdTest,
)
def _rmsnorm(x, w, eps=1e-6):

View File

@ -7,7 +7,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import FFTC2COp
from workloads.base import FixtureBase
from workloads.ops.fft import FFTTest
from workloads.fft import FFTTest
class _FFTTestBaseline(FFTTest):

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import FP8LightingIndexerOp
from workloads.ops.fp8_lighting_indexer import FP8LightingIndexerTest
from workloads.fp8_lighting_indexer import FP8LightingIndexerTest
class _FP8LightingIndexerTestBaseline(FP8LightingIndexerTest):

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import FP8QuantOp
from workloads.ops.fp8_quant import FP8QuantTest
from workloads.fp8_quant import FP8QuantTest
class _FP8QuantTestBaseline(FP8QuantTest):

View File

@ -7,7 +7,7 @@ import torch.nn.functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.fused_add_layer_norm import FusedAddLayerNormFwdOp
from workloads.ops.fused_add_layer_norm import FusedAddLayerNormTest
from workloads.fused_add_layer_norm import FusedAddLayerNormTest
_OP_NAME = "FusedAddLayerNormFwdOp"

View File

@ -6,7 +6,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.fused_add_rmsnorm import FusedAddRMSNormFwdOp
from workloads.ops.fused_add_rmsnorm import FusedAddRmsNormTest
from workloads.fused_add_rmsnorm import FusedAddRmsNormTest
_OP_NAME = "FusedAddRMSNormFwdOp"

View File

@ -21,7 +21,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GatedDeltaNetBwdOp, GatedDeltaNetFwdOp, GatedDeltaNetOp
from workloads.base import FixtureBase
from workloads.ops.gated_deltanet_fwd import GatedDeltaNetFwdTest
from workloads.gated_deltanet import GatedDeltaNetFwdTest
def _differentiable_fwd(q, k, v, g_raw, beta, chunk_size):

View File

@ -6,7 +6,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GatedDeltaNetDecodeOp
from workloads.base import FixtureBase
from workloads.ops.gated_deltanet_recurrence import GatedDeltaNetDecodeTest
from workloads.gated_deltanet import GatedDeltaNetDecodeTest
def gated_deltanet_decode_torch(

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GemmOp
from workloads.ops.gemm import GemmTest
from workloads.gemm import GemmTest
class _GemmTestBaseline(GemmTest):

View File

@ -13,7 +13,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GLADecodeOp
from workloads.base import FixtureBase
from workloads.ops.gla_recurrence import GLADecodeTest
from workloads.gla import GLADecodeTest
def gla_decode_torch(

View File

@ -8,7 +8,7 @@ import torch.nn.functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.group_norm import GroupNormFwdOp
from workloads.ops.group_norm import GroupNormTest
from workloads.group_norm import GroupNormTest
_OP_NAME = "GroupNormFwdOp"

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import GroupedGemmOp
from workloads.ops.grouped_gemm import (
from workloads.grouped_gemm import (
GroupedGemmCompleteTest,
GroupedGemmTest,
)

View File

@ -8,7 +8,7 @@ import torch.nn.functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.instance_norm import InstanceNormFwdOp
from workloads.ops.instance_norm import InstanceNormTest
from workloads.instance_norm import InstanceNormTest
_OP_NAME = "InstanceNormFwdOp"

View File

@ -7,7 +7,7 @@ import torch.nn.functional as F
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.layer_norm import LayerNormFwdOp
from workloads.ops.layer_norm import LayerNormTest
from workloads.layer_norm import LayerNormTest
_OP_NAME = "LayerNormFwdOp"

View File

@ -11,7 +11,7 @@ from benchmarks.benchmark import BenchmarkReport, ManifestBenchmark, workloads_t
from tileops.ops.reduction.all_op import AllFwdOp
from tileops.ops.reduction.any_op import AnyFwdOp
from tileops.ops.reduction.count_nonzero import CountNonzeroFwdOp
from workloads.ops.logical_reduce import AllTest, AnyTest, CountNonzeroTest
from workloads.logical_reduce import AllTest, AnyTest, CountNonzeroTest
# ===================================================================
# Op name constants

View File

@ -9,11 +9,18 @@ from tileops.ops.ssd_chunk_scan import SsdChunkScanFwdOp
from tileops.ops.ssd_chunk_state import SsdChunkStateFwdOp
from tileops.ops.ssd_decode import SsdDecodeOp
from tileops.ops.ssd_state_passing import SsdStatePassingFwdOp
from workloads.ops.da_cumsum_fwd import DaCumsumFwdFixture, DaCumsumFwdTest
from workloads.ops.ssd_chunk_scan_fwd import SsdChunkScanFwdFixture, SsdChunkScanFwdTest
from workloads.ops.ssd_chunk_state_fwd import SsdChunkStateFwdFixture, SsdChunkStateFwdTest
from workloads.ops.ssd_decode import SsdDecodeFixture, SsdDecodeTest
from workloads.ops.ssd_state_passing_fwd import SsdStatePassingFwdFixture, SsdStatePassingFwdTest
from workloads.mamba import (
DaCumsumFwdFixture,
DaCumsumFwdTest,
SsdChunkScanFwdFixture,
SsdChunkScanFwdTest,
SsdChunkStateFwdFixture,
SsdChunkStateFwdTest,
SsdDecodeFixture,
SsdDecodeTest,
SsdStatePassingFwdFixture,
SsdStatePassingFwdTest,
)
def da_cumsum_fwd_ref(

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import MHCPostOp
from workloads.ops.mhc_post import MHCPostTest
from workloads.mhc import MHCPostTest
class _MHCPostTestBaseline(MHCPostTest):

View File

@ -6,7 +6,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import MHCPreOp
from workloads.ops.mhc_pre import MHCPreTest
from workloads.mhc import MHCPreTest
class _MHCPreTestBaseline(MHCPreTest):

View File

@ -29,7 +29,7 @@ except ImportError:
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops.moe import FusedTopKOp
from workloads.base import FixtureBase
from workloads.ops.moe_fused_topk import FusedTopKTest
from workloads.moe import FusedTopKTest
def fused_topk_torch(

View File

@ -32,7 +32,7 @@ except ImportError:
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.moe import MoePermutePaddedFwdOp
from workloads.ops.moe_permute import MoePermuteTest
from workloads.moe import MoePermuteTest
_OP_NAME = "MoePermutePaddedFwdOp"

View File

@ -28,7 +28,7 @@ except ImportError:
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.moe import MoePermuteAlignFwdOp
from workloads.ops.moe_permute_align import MoePermuteAlignTest
from workloads.moe import MoePermuteAlignTest
_OP_NAME = "MoePermuteAlignFwdOp"

View File

@ -33,7 +33,7 @@ except ImportError:
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.moe import MoeUnpermuteFwdOp
from workloads.ops.moe_unpermute import MoeUnpermuteTest
from workloads.moe import MoeUnpermuteTest
_OP_NAME = "MoeUnpermuteFwdOp"

View File

@ -18,7 +18,7 @@ from tileops.ops.reduction.reduce import (
VarFwdOp,
VarMeanFwdOp,
)
from workloads.ops.reduce import (
from workloads.reduce import (
AmaxTest,
AminTest,
MeanTest,

View File

@ -6,7 +6,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.manifest import eval_roofline, load_workloads
from tileops.ops.norm.rms_norm import RMSNormFwdOp
from workloads.ops.rms_norm import RmsNormTest
from workloads.rms_norm import RmsNormTest
class _RmsNormTestBaseline(RmsNormTest):

View File

@ -12,7 +12,7 @@ from benchmarks.benchmark import BenchmarkReport, ManifestBenchmark, workloads_t
from tileops.ops.reduction.log_softmax import LogSoftmaxFwdOp
from tileops.ops.reduction.logsumexp import LogSumExpFwdOp
from tileops.ops.reduction.softmax import SoftmaxFwdOp
from workloads.ops.softmax import (
from workloads.softmax import (
LogSoftmaxTest,
LogSumExpTest,
SoftmaxTest,

View File

@ -5,7 +5,7 @@ import torch
from benchmarks.benchmark import BenchmarkBase, BenchmarkReport
from tileops.ops import TopkSelectorOp
from workloads.ops.topk_selector import TopkSelectorTest
from workloads.topk_selector import TopkSelectorTest
class _TopkSelectorTestBaseline(TopkSelectorTest):

View File

@ -11,7 +11,7 @@ from benchmarks.benchmark import BenchmarkReport, ManifestBenchmark, workloads_t
from tileops.ops.reduction.inf_norm import InfNormFwdOp
from tileops.ops.reduction.l1_norm import L1NormFwdOp
from tileops.ops.reduction.l2_norm import L2NormFwdOp
from workloads.ops.vector_norm import InfNormTest, L1NormTest, L2NormTest
from workloads.vector_norm import InfNormTest, L1NormTest, L2NormTest
# ===================================================================
# Op name constants

View File

@ -15,7 +15,7 @@ Tests and benchmarks are separated by concern: `pytest tests/` validates correct
## Test/Benchmark Pattern
```python
# workloads/ops/mha.py
# workloads/attention/mha.py
class MhaFwdTest(WorkloadBase):
def __init__(self, batch, heads, seq_len, dim, causal, dtype): ...
def gen_inputs(self): ...
@ -23,7 +23,7 @@ class MhaFwdTest(WorkloadBase):
# tests/ops/test_mha.py
from tileops.ops import MhaFwdOp
from workloads.ops.mha import MhaFwdTest
from workloads.attention.mha import MhaFwdTest
class MhaFwdTestCase(MhaFwdTest, TestBase):
@ -43,7 +43,7 @@ def test_mha_fwd(batch, seq_len, heads, dim, causal, dtype, tune):
# benchmarks/ops/bench_mha.py
from tileops.ops import MhaFwdOp
from workloads.ops.mha import MhaFwdTest # import workload, NOT test
from workloads.attention.mha import MhaFwdTest # import workload, NOT test
class MhaFwdBenchmark(BenchmarkBase):
@ -151,7 +151,7 @@ python scripts/test_node_delta.py --base origin/release # different base branc
### File checklist
1. **Workload class** in `workloads/ops/` — subclass `WorkloadBase`, implement `gen_inputs()`.
1. **Workload class** in `workloads/` — subclass `WorkloadBase`, implement `gen_inputs()`.
1. **Fixture class** — subclass `FixtureBase`, define `PARAMS` with `smoke`/`full` marks.
1. **Test class** in `tests/ops/test_<op>.py` — inherit `(MyWorkload, TestBase)`, implement `ref_program()` locally.
1. **Test function**`@YourFixture` decorated, call `test.check(op, *test.gen_inputs())`.
@ -177,7 +177,7 @@ TestBase (tests/test_base.py, inherits WorkloadBase)
### File checklist
1. **Workload class** in `workloads/ops/` — reuse the `WorkloadBase` subclass from the test.
1. **Workload class** in `workloads/` — reuse the `WorkloadBase` subclass from the test.
1. **Fixture class** — reuse the `FixtureBase` subclass from the test.
1. **Benchmark class** in `benchmarks/ops/bench_<op>.py` — subclass `BenchmarkBase`, implement `calculate_flops()` and `calculate_memory()` (return `None` if not applicable).
1. **Benchmark function**`@YourFixture` decorated, construct workload + benchmark, call `inputs = workload.gen_inputs()`, then `bm.profile(op, *inputs)` and `BenchmarkReport.record(op, locals(), result, tag="tileops")`.

View File

@ -4,7 +4,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import DeepSeekSparseAttentionDecodeWithKVCacheFwdOp
from workloads.ops.deepseek_dsa_decode import DsaDecodeTest as _DsaDecodeTestWorkload
from workloads.attention.deepseek_dsa_decode import DsaDecodeTest as _DsaDecodeTestWorkload
class DsaDecodeTest(_DsaDecodeTestWorkload, TestBase):

View File

@ -6,7 +6,7 @@ from einops import einsum, rearrange
from tests.test_base import FixtureBase, TestBase
from tileops.ops import MultiHeadLatentAttentionDecodeWithKVCacheFwdOp
from workloads.ops.deepseek_mla_decode import MlaDecodeTest as _MlaDecodeTestWorkload
from workloads.attention.deepseek_mla_decode import MlaDecodeTest as _MlaDecodeTestWorkload
class MlaDecodeTest(_MlaDecodeTestWorkload, TestBase):

View File

@ -6,7 +6,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import NSAFwdVarlenOp
from workloads.ops.deepseek_nsa_fwd import NsaFwdTest as _NsaFwdTestWorkload
from workloads.attention.deepseek_nsa import NsaFwdTest as _NsaFwdTestWorkload
class NsaFwdTest(_NsaFwdTestWorkload, TestBase):

View File

@ -6,8 +6,8 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import NSACmpFwdVarlenOp
from workloads.attention.deepseek_nsa_cmp import NsaCmpFwdTest as _NsaCmpFwdTestWorkload
from workloads.nsa_utils import prepare_chunk_offsets
from workloads.ops.deepseek_nsa_cmp_fwd import NsaCmpFwdTest as _NsaCmpFwdTestWorkload
def _parallel_nsa_compression_fwd_pytorch(test, q, k_cmp, v_cmp, block_size, scale, offsets):

View File

@ -4,7 +4,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import NSATopkVarlenOp
from workloads.ops.deepseek_nsa_topk import NsaTopkTest as _NsaTopkTestWorkload
from workloads.attention.deepseek_nsa_topk import NsaTopkTest as _NsaTopkTestWorkload
def _nsa_topk_torch(test, q, k_cmp, lse, block_counts, block_size, scale,

View File

@ -8,8 +8,8 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GroupedQueryAttentionBwdOp, GroupedQueryAttentionFwdOp
from workloads.ops.gqa import GqaBwdTest as _GqaBwdTestWorkload
from workloads.ops.gqa import GqaFwdTest as _GqaFwdTestWorkload
from workloads.attention.gqa import GqaBwdTest as _GqaBwdTestWorkload
from workloads.attention.gqa import GqaFwdTest as _GqaFwdTestWorkload
class GqaBwdTest(_GqaBwdTestWorkload, TestBase):

View File

@ -6,7 +6,7 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GroupedQueryAttentionDecodeWithKVCacheFwdOp
from workloads.ops.gqa_decode import GqaDecodeTest as _GqaDecodeTestWorkload
from workloads.attention.gqa_decode import GqaDecodeTest as _GqaDecodeTestWorkload
class GqaDecodeTest(_GqaDecodeTestWorkload, TestBase):

View File

@ -10,7 +10,9 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GroupedQueryAttentionDecodePagedWithKVCacheFwdOp
from workloads.ops.gqa_decode_paged import GqaDecodePagedTest as _GqaDecodePagedTestWorkload
from workloads.attention.gqa_decode_paged import (
GqaDecodePagedTest as _GqaDecodePagedTestWorkload,
)
class GqaDecodePagedTest(_GqaDecodePagedTestWorkload, TestBase):

View File

@ -5,7 +5,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GqaSlidingWindowFwdOp
from workloads.ops.gqa_sliding_window_fwd import (
from workloads.attention.gqa_sliding_window import (
GqaSlidingWindowFwdTest as _GqaSlidingWindowFwdTestWorkload,
)

View File

@ -6,7 +6,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GqaSlidingWindowVarlenFwdOp
from workloads.ops.gqa_sliding_window_varlen_fwd import (
from workloads.attention.gqa_sliding_window_varlen import (
GqaSlidingWindowVarlenFwdTest as _GqaSlidingWindowVarlenFwdTestWorkload,
)

View File

@ -5,8 +5,8 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import MeanPoolingForwardOp
from workloads.attention.mean_pooling import MeanPoolingTest as _MeanPoolingTestWorkload
from workloads.nsa_utils import prepare_chunk_indices
from workloads.ops.mean_pooling_ops import MeanPoolingTest as _MeanPoolingTestWorkload
class MeanPoolingTest(_MeanPoolingTestWorkload, TestBase):

View File

@ -8,8 +8,8 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from tests.test_base import FixtureBase, TestBase
from tileops.ops import MultiHeadAttentionBwdOp, MultiHeadAttentionFwdOp
from workloads.ops.mha import MhaBwdTest as _MhaBwdTestWorkload
from workloads.ops.mha import MhaFwdTest as _MhaFwdTestWorkload
from workloads.attention.mha import MhaBwdTest as _MhaBwdTestWorkload
from workloads.attention.mha import MhaFwdTest as _MhaFwdTestWorkload
class MhaBwdTest(_MhaBwdTestWorkload, TestBase):

View File

@ -6,7 +6,7 @@ from torch.nn.attention import SDPBackend, sdpa_kernel
from tests.test_base import FixtureBase, TestBase
from tileops.ops import MultiHeadAttentionDecodeWithKVCacheFwdOp
from workloads.ops.mha_decode import MhaDecodeTest as _MhaDecodeTestWorkload
from workloads.attention.mha_decode import MhaDecodeTest as _MhaDecodeTestWorkload
class MhaDecodeTest(_MhaDecodeTestWorkload, TestBase):

View File

@ -9,7 +9,9 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops import MultiHeadAttentionDecodePagedWithKVCacheFwdOp
from workloads.ops.mha_decode_paged import MhaDecodePagedTest as _MhaDecodePagedTestWorkload
from workloads.attention.mha_decode_paged import (
MhaDecodePagedTest as _MhaDecodePagedTestWorkload,
)
class MhaDecodePagedTest(_MhaDecodePagedTestWorkload, TestBase):

View File

@ -9,7 +9,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.elementwise import ReluOp
from workloads.ops.activation import ReluTest as _ReluTestWorkload
from workloads.activation import ReluTest as _ReluTestWorkload
class ReluTest(_ReluTestWorkload, TestBase):

View File

@ -4,7 +4,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.ada_layer_norm import AdaLayerNormFwdOp
from workloads.ops.ada_layer_norm import AdaLayerNormTest as _AdaLayerNormTestWorkload
from workloads.ada_layer_norm import AdaLayerNormTest as _AdaLayerNormTestWorkload
class AdaLayerNormTest(_AdaLayerNormTestWorkload, TestBase):

View File

@ -4,7 +4,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.ada_layer_norm_zero import AdaLayerNormZeroFwdOp
from workloads.ops.ada_layer_norm_zero import AdaLayerNormZeroTest as _AdaLayerNormZeroTestWorkload
from workloads.ada_layer_norm_zero import AdaLayerNormZeroTest as _AdaLayerNormZeroTestWorkload
class AdaLayerNormZeroTest(_AdaLayerNormZeroTestWorkload, TestBase):

View File

@ -9,7 +9,7 @@ import pytest
import torch
from tests.test_base import FixtureBase, TestBase
from workloads.ops.argreduce import ArgmaxTest as _ArgmaxWorkload
from workloads.argreduce import ArgmaxTest as _ArgmaxWorkload
# ---------------------------------------------------------------------------
# Fixtures

View File

@ -13,10 +13,10 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.batch_norm import BatchNormBwdOp, BatchNormFwdOp
from workloads.ops.batch_norm import (
from workloads.batch_norm import (
BatchNormBwdTest as _BatchNormBwdTestWorkload,
)
from workloads.ops.batch_norm import (
from workloads.batch_norm import (
BatchNormFwdTest as _BatchNormFwdTestWorkload,
)

View File

@ -22,7 +22,7 @@ from tileops.ops.elementwise import (
SubOp,
coalesce_broadcast_dims,
)
from workloads.ops.binary_arith import AddSameShapeTest as _AddSameShapeTestWorkload
from workloads.binary_arith import AddSameShapeTest as _AddSameShapeTestWorkload
class AddSameShapeTest(_AddSameShapeTestWorkload, TestBase):

View File

@ -4,7 +4,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import DeltaNetFwdOp
from workloads.ops.deltanet_fwd import DeltaNetFwdTest as _DeltaNetFwdTestWorkload
from workloads.deltanet import DeltaNetFwdTest as _DeltaNetFwdTestWorkload
def compute_w_u_torch(Aw, Au, k, v, beta, chunk_size):

View File

@ -6,7 +6,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import DeltaNetDecodeOp
from workloads.ops.deltanet_recurrence import DeltaNetDecodeTest as _DeltaNetDecodeTestWorkload
from workloads.deltanet import DeltaNetDecodeTest as _DeltaNetDecodeTestWorkload
def deltanet_decode_torch(

View File

@ -5,10 +5,16 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.engram import EngramGateConvBwdOp, EngramGateConvFwdOp
from tileops.ops.engram_decode import EngramDecodeOp
from workloads.ops.engram_bwd import EngramGateConvBwdTest as _EngramGateConvBwdTestWorkload
from workloads.ops.engram_decode import EngramDecodeTest as _EngramDecodeTestWorkload
from workloads.ops.engram_fwd import CONV_KERNEL_SIZE
from workloads.ops.engram_fwd import (
from workloads.engram import (
CONV_KERNEL_SIZE,
)
from workloads.engram import (
EngramDecodeTest as _EngramDecodeTestWorkload,
)
from workloads.engram import (
EngramGateConvBwdTest as _EngramGateConvBwdTestWorkload,
)
from workloads.engram import (
EngramGateConvFwdTest as _EngramGateConvFwdTestWorkload,
)

View File

@ -4,7 +4,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import FFTC2COp
from workloads.ops.fft import FFTTest as _FFTTestWorkload
from workloads.fft import FFTTest as _FFTTestWorkload
class FFTTest(_FFTTestWorkload, TestBase):

View File

@ -5,7 +5,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import FP8LightingIndexerOp
from workloads.ops.fp8_lighting_indexer import (
from workloads.fp8_lighting_indexer import (
FP8LightingIndexerTest as _FP8LightingIndexerTestWorkload,
)

View File

@ -7,7 +7,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops import FP8QuantOp
from workloads.ops.fp8_quant import FP8QuantTest as _FP8QuantTestWorkload
from workloads.fp8_quant import FP8QuantTest as _FP8QuantTestWorkload
class FP8QuantTest(_FP8QuantTestWorkload, TestBase):

View File

@ -4,7 +4,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.fused_add_layer_norm import FusedAddLayerNormFwdOp
from workloads.ops.fused_add_layer_norm import (
from workloads.fused_add_layer_norm import (
FusedAddLayerNormTest as _FusedAddLayerNormTestWorkload,
)

View File

@ -3,7 +3,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.fused_add_rmsnorm import FusedAddRMSNormFwdOp
from workloads.ops.fused_add_rmsnorm import FusedAddRmsNormTest as _FusedAddRmsNormTestWorkload
from workloads.fused_add_rmsnorm import FusedAddRmsNormTest as _FusedAddRmsNormTestWorkload
class FusedAddRmsNormTest(_FusedAddRmsNormTestWorkload, TestBase):

View File

@ -4,7 +4,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GatedDeltaNetFwdOp
from workloads.ops.gated_deltanet_fwd import (
from workloads.gated_deltanet import (
GatedDeltaNetFwdTest as _GatedDeltaNetFwdTestWorkload,
)

View File

@ -6,7 +6,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GatedDeltaNetDecodeOp
from workloads.ops.gated_deltanet_recurrence import (
from workloads.gated_deltanet import (
GatedDeltaNetDecodeTest as _GatedDeltaNetDecodeTestWorkload,
)

View File

@ -4,7 +4,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GemmOp
from workloads.ops.gemm import GemmTest as _GemmTestWorkload
from workloads.gemm import GemmTest as _GemmTestWorkload
class GemmTest(_GemmTestWorkload, TestBase):

View File

@ -6,7 +6,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import GLADecodeOp
from workloads.ops.gla_recurrence import GLADecodeTest as _GLADecodeTestWorkload
from workloads.gla import GLADecodeTest as _GLADecodeTestWorkload
def gla_decode_torch(

View File

@ -4,7 +4,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.group_norm import GroupNormFwdOp
from workloads.ops.group_norm import GroupNormTest as _GroupNormTestWorkload
from workloads.group_norm import GroupNormTest as _GroupNormTestWorkload
class GroupNormTest(_GroupNormTestWorkload, TestBase):

View File

@ -4,7 +4,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops.grouped_gemm import GroupedGemmOp
from workloads.ops.grouped_gemm import (
from workloads.grouped_gemm import (
GroupedGemmTest as _GroupedGemmTestWorkload,
)

View File

@ -4,7 +4,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.instance_norm import InstanceNormFwdOp
from workloads.ops.instance_norm import InstanceNormTest as _InstanceNormTestWorkload
from workloads.instance_norm import InstanceNormTest as _InstanceNormTestWorkload
class InstanceNormTest(_InstanceNormTestWorkload, TestBase):

View File

@ -4,7 +4,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.layer_norm import LayerNormFwdOp
from workloads.ops.layer_norm import LayerNormTest as _LayerNormTestWorkload
from workloads.layer_norm import LayerNormTest as _LayerNormTestWorkload
class LayerNormTest(_LayerNormTestWorkload, TestBase):

View File

@ -10,7 +10,7 @@ import pytest
import torch
from tests.test_base import FixtureBase, TestBase
from workloads.ops.logical_reduce import AnyTest as _AnyWorkload
from workloads.logical_reduce import AnyTest as _AnyWorkload
# ---------------------------------------------------------------------------
# Fixtures

View File

@ -7,16 +7,18 @@ from tileops.ops.ssd_chunk_scan import SsdChunkScanFwdOp
from tileops.ops.ssd_chunk_state import SsdChunkStateFwdOp
from tileops.ops.ssd_decode import SsdDecodeOp
from tileops.ops.ssd_state_passing import SsdStatePassingFwdOp
from workloads.ops.da_cumsum_fwd import DaCumsumFwdFixture
from workloads.ops.da_cumsum_fwd import DaCumsumFwdTest as _DaCumsumFwdTestWorkload
from workloads.ops.ssd_chunk_scan_fwd import SsdChunkScanFwdFixture
from workloads.ops.ssd_chunk_scan_fwd import SsdChunkScanFwdTest as _SsdChunkScanFwdTestWorkload
from workloads.ops.ssd_chunk_state_fwd import SsdChunkStateFwdFixture
from workloads.ops.ssd_chunk_state_fwd import SsdChunkStateFwdTest as _SsdChunkStateFwdTestWorkload
from workloads.ops.ssd_decode import SsdDecodeFixture
from workloads.ops.ssd_decode import SsdDecodeTest as _SsdDecodeTestWorkload
from workloads.ops.ssd_state_passing_fwd import SsdStatePassingFwdFixture
from workloads.ops.ssd_state_passing_fwd import (
from workloads.mamba import (
DaCumsumFwdFixture,
SsdChunkScanFwdFixture,
SsdChunkStateFwdFixture,
SsdDecodeFixture,
SsdStatePassingFwdFixture,
)
from workloads.mamba import DaCumsumFwdTest as _DaCumsumFwdTestWorkload
from workloads.mamba import SsdChunkScanFwdTest as _SsdChunkScanFwdTestWorkload
from workloads.mamba import SsdChunkStateFwdTest as _SsdChunkStateFwdTestWorkload
from workloads.mamba import SsdDecodeTest as _SsdDecodeTestWorkload
from workloads.mamba import (
SsdStatePassingFwdTest as _SsdStatePassingFwdTestWorkload,
)

View File

@ -7,7 +7,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops import MHCPostOp
from workloads.ops.mhc_post import MHCPostTest as _MHCPostTestWorkload
from workloads.mhc import MHCPostTest as _MHCPostTestWorkload
class MHCPostTest(_MHCPostTestWorkload, TestBase):

View File

@ -10,7 +10,7 @@ import torch.nn.functional as F
from tests.test_base import FixtureBase, TestBase
from tileops.ops import MHCPreOp
from workloads.ops.mhc_pre import MHCPreTest as _MHCPreTestWorkload
from workloads.mhc import MHCPreTest as _MHCPreTestWorkload
class MHCPreTest(_MHCPreTestWorkload, TestBase):

View File

@ -16,7 +16,7 @@ import torch
from tests.test_base import FixtureBase
from tileops.ops.moe import FusedTopKOp
from workloads.ops.moe_fused_topk import FusedTopKTest
from workloads.moe import FusedTopKTest
def fused_topk_torch(

View File

@ -16,7 +16,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops.moe import MoePermutePaddedFwdOp
from workloads.ops.moe_permute import MoePermuteTest as _MoePermuteTestWorkload
from workloads.moe import MoePermuteTest as _MoePermuteTestWorkload
def _ref_moe_permute(

View File

@ -15,7 +15,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops.moe import MoePermuteAlignFwdOp
from workloads.ops.moe_permute_align import MoePermuteAlignTest as _MoePermuteAlignTestWorkload
from workloads.moe import MoePermuteAlignTest as _MoePermuteAlignTestWorkload
def _ref_permute_align(

View File

@ -18,7 +18,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops.moe import MoeUnpermuteFwdOp
from workloads.ops.moe_unpermute import MoeUnpermuteTest as _MoeUnpermuteTestWorkload
from workloads.moe import MoeUnpermuteTest as _MoeUnpermuteTestWorkload
def _ref_moe_unpermute(

View File

@ -8,13 +8,13 @@ import pytest
import torch
from tests.test_base import FixtureBase, TestBase
from workloads.ops.reduce import (
from workloads.reduce import (
ProdTest as _ProdTest,
)
from workloads.ops.reduce import (
from workloads.reduce import (
StdTest as _StdTest,
)
from workloads.ops.reduce import (
from workloads.reduce import (
SumTest as _SumTest,
)

View File

@ -3,7 +3,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops.norm.rms_norm import RMSNormFwdOp
from workloads.ops.rms_norm import RmsNormTest as _RmsNormTestWorkload
from workloads.rms_norm import RmsNormTest as _RmsNormTestWorkload
class RmsNormTest(_RmsNormTestWorkload, TestBase):

View File

@ -21,9 +21,9 @@ from tests.test_base import FixtureBase, TestBase
from tileops.ops.reduction.log_softmax import LogSoftmaxFwdOp
from tileops.ops.reduction.logsumexp import LogSumExpFwdOp
from tileops.ops.reduction.softmax import SoftmaxFwdOp
from workloads.ops.softmax import LogSoftmaxTest as _LogSoftmaxTestWorkload
from workloads.ops.softmax import LogSumExpTest as _LogSumExpTestWorkload
from workloads.ops.softmax import SoftmaxTest as _SoftmaxTestWorkload
from workloads.softmax import LogSoftmaxTest as _LogSoftmaxTestWorkload
from workloads.softmax import LogSumExpTest as _LogSumExpTestWorkload
from workloads.softmax import SoftmaxTest as _SoftmaxTestWorkload
# ---------------------------------------------------------------------------
# Tolerances (from docs/testing.md)

View File

@ -5,7 +5,7 @@ import torch
from tests.test_base import FixtureBase, TestBase
from tileops.ops import TopkSelectorOp
from tileops.utils import str2dtype
from workloads.ops.topk_selector import TopkSelectorTest as _TopkSelectorTestWorkload
from workloads.topk_selector import TopkSelectorTest as _TopkSelectorTestWorkload
class TopkSelectorTest(_TopkSelectorTestWorkload, TestBase):

View File

@ -9,7 +9,7 @@ import pytest
import torch
from tests.test_base import FixtureBase, TestBase, allclose_compare
from workloads.ops.vector_norm import L1NormTest as _L1NormWorkload
from workloads.vector_norm import L1NormTest as _L1NormWorkload
# ---------------------------------------------------------------------------
# Fixtures

Some files were not shown because too many files have changed in this diff Show More