forked from ccf-ai-infra/TileOPs-Metax
[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:
parent
b77f40c762
commit
6e0b507b42
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ from tileops.ops.reduction.reduce import (
|
|||
VarFwdOp,
|
||||
VarMeanFwdOp,
|
||||
)
|
||||
from workloads.ops.reduce import (
|
||||
from workloads.reduce import (
|
||||
AmaxTest,
|
||||
AminTest,
|
||||
MeanTest,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")`.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Reference in New Issue