diff --git a/docs/fe-oss-apis/attention/sdpa_bwd_d256.md b/docs/fe-oss-apis/attention/sdpa_bwd_d256.md index e183af066..23a24dafe 100644 --- a/docs/fe-oss-apis/attention/sdpa_bwd_d256.md +++ b/docs/fe-oss-apis/attention/sdpa_bwd_d256.md @@ -275,6 +275,6 @@ Tuple unpacking order is: `(dq_tensor, dk_tensor, dv_tensor)`. For runnable examples and reference-comparison checks, see: -- `test/python/fe_api/test_sdpa_bwd.py` -- `test/python/fe_api/test_sdpa_bwd_utils.py` +- `test/python/fe_api/sdpa/test_sdpa_bwd.py` +- `test/python/fe_api/sdpa/test_sdpa_bwd_utils.py` diff --git a/docs/fe-oss-apis/bsa.md b/docs/fe-oss-apis/bsa.md index 09af7649e..b39643365 100644 --- a/docs/fe-oss-apis/bsa.md +++ b/docs/fe-oss-apis/bsa.md @@ -214,7 +214,7 @@ The current public surface consists of allocating function wrappers under lifecycle for BSA. Correctness tests and FP32 references are under -`test/python/fe_api/block_sparse_attention`. +`test/python/fe_api/bsa`. ## Source provenance diff --git a/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_dswiglu.md b/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_dswiglu.md index 1dc18c5d3..b80675432 100644 --- a/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_dswiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_dswiglu.md @@ -465,5 +465,5 @@ Returns a `TupleDict` - a dictionary-like object that also supports tuple unpack ## Usage Examples For runnable examples and validation, see: -- `test/python/fe_api/test_discrete_grouped_gemm_dswiglu.py` -- `test/python/fe_api/test_discrete_grouped_gemm_dswiglu_utils.py` +- `test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu.py` +- `test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_swiglu.md b/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_swiglu.md index 043c720fb..d89841c24 100644 --- a/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_swiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_swiglu.md @@ -424,5 +424,5 @@ Returns a `TupleDict` - a dictionary-like object that also supports tuple unpack ## Usage Examples For runnable examples and validation, see: -- `test/python/fe_api/test_discrete_grouped_gemm_swiglu.py` -- `test/python/fe_api/test_discrete_grouped_gemm_swiglu_utils.py` +- `test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu.py` +- `test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/gemm_amax.md b/docs/fe-oss-apis/gemm_fusions/gemm_amax.md index 12f14f11e..a39cfb26a 100644 --- a/docs/fe-oss-apis/gemm_fusions/gemm_amax.md +++ b/docs/fe-oss-apis/gemm_fusions/gemm_amax.md @@ -204,4 +204,4 @@ Tuple unpacking order is: `(c_tensor, amax_tensor)`. ## Usage examples -For usage examples, see test cases in `test/python/fe_api/test_gemm_amax.py` +For usage examples, see test cases in `test/python/fe_api/gemm/test_gemm_amax.py` diff --git a/docs/fe-oss-apis/gemm_fusions/gemm_dsrelu.md b/docs/fe-oss-apis/gemm_fusions/gemm_dsrelu.md index a18656744..7e7f17ab1 100644 --- a/docs/fe-oss-apis/gemm_fusions/gemm_dsrelu.md +++ b/docs/fe-oss-apis/gemm_fusions/gemm_dsrelu.md @@ -249,5 +249,5 @@ Tuple unpacking order is: `(d_tensor, dprob_tensor, amax_tensor, sfd_tensor)`. For end-to-end usage and regression coverage, see: -- `test/python/fe_api/test_gemm_dsrelu.py` -- `test/python/fe_api/test_gemm_dsrelu_utils.py` +- `test/python/fe_api/gemm/test_gemm_dsrelu.py` +- `test/python/fe_api/gemm/test_gemm_dsrelu_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/gemm_srelu.md b/docs/fe-oss-apis/gemm_fusions/gemm_srelu.md index 963165e3f..29f490644 100644 --- a/docs/fe-oss-apis/gemm_fusions/gemm_srelu.md +++ b/docs/fe-oss-apis/gemm_fusions/gemm_srelu.md @@ -234,5 +234,5 @@ Tuple unpacking order is: `(c_tensor, d_tensor, amax_tensor, sfd_tensor)`. For end-to-end usage and regression coverage, see: -- `test/python/fe_api/test_gemm_srelu.py` -- `test/python/fe_api/test_gemm_srelu_utils.py` +- `test/python/fe_api/gemm/test_gemm_srelu.py` +- `test/python/fe_api/gemm/test_gemm_srelu_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/gemm_swiglu.md b/docs/fe-oss-apis/gemm_fusions/gemm_swiglu.md index e03c041e4..ce18625e3 100644 --- a/docs/fe-oss-apis/gemm_fusions/gemm_swiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/gemm_swiglu.md @@ -350,4 +350,4 @@ Additional constraints: ## Usage examples -For usage examples, see test cases in `test/python/fe_api/test_gemm_swiglu.py` +For usage examples, see test cases in `test/python/fe_api/gemm/test_gemm_swiglu.py` diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dglu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dglu.md index 43cf67715..c2a0f9bdd 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dglu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dglu.md @@ -396,4 +396,4 @@ Returns a `TupleDict` (dictionary + tuple unpacking): ## Usage Examples -For usage examples, see test cases in `test/python/fe_api/test_grouped_gemm_dglu.py` (dense mode, unified API) and `test/python/fe_api/test_discrete_grouped_gemm_dswiglu.py` (discrete mode). +For usage examples, see test cases in `test/python/fe_api/grouped_gemm/test_grouped_gemm_dglu.py` (dense mode, unified API) and `test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu.py` (discrete mode). diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dsrelu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dsrelu.md index 128276284..e78a5f0a8 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dsrelu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dsrelu.md @@ -314,5 +314,5 @@ Tuple unpacking order is: `(d_row_tensor, d_col_tensor, dprob_tensor, dbias_tens For end-to-end usage and regression coverage, see: -- `test/python/fe_api/test_grouped_gemm_dsrelu.py` -- `test/python/fe_api/test_grouped_gemm_dsrelu_utils.py` +- `test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu.py` +- `test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md index c7f1de2e1..cf7b40fa8 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu.md @@ -410,4 +410,4 @@ Returns a `TupleDict` - a dictionary-like object that also supports tuple unpack ## Usage Examples -For usage examples, see test cases in `test/python/fe_api/test_grouped_gemm_dswiglu.py` + `test/python/fe_api/test_grouped_gemm_dswiglu_utils.py` +For usage examples, see test cases in `test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu.py` + `test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_glu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_glu.md index 2061c8be4..320912606 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_glu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_glu.md @@ -389,4 +389,4 @@ Returns a `TupleDict` (dictionary + tuple unpacking): ## Usage Examples -For usage examples, see test cases in `test/python/fe_api/test_grouped_gemm_glu.py` (dense mode, unified API) and `test/python/fe_api/test_discrete_grouped_gemm_swiglu.py` (discrete mode). +For usage examples, see test cases in `test/python/fe_api/grouped_gemm/test_grouped_gemm_glu.py` (dense mode, unified API) and `test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu.py` (discrete mode). diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant.md index 3d987fae4..550854e6f 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant.md @@ -372,4 +372,4 @@ Returns a `TupleDict` - a dictionary-like object that also supports tuple unpack ## Usage Examples -For usage examples, see test cases in `test/python/fe_api/test_grouped_gemm_quant.py` + `test/python/fe_api/test_grouped_gemm_quant_utils.py` +For usage examples, see test cases in `test/python/fe_api/grouped_gemm/test_grouped_gemm_quant.py` + `test/python/fe_api/grouped_gemm/test_grouped_gemm_quant_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant_unified.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant_unified.md index daec4e61d..ec882afbf 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant_unified.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_quant_unified.md @@ -244,4 +244,4 @@ Returns `TupleDict`: `d_tensor`, `d_col_tensor` (optional; `None` for `bfloat16` ## Usage Examples -For usage examples, see `test/python/fe_api/test_grouped_gemm_quant.py` + `test/python/fe_api/test_grouped_gemm_quant_utils.py` (dense and discrete unified API coverage) +For usage examples, see `test/python/fe_api/grouped_gemm/test_grouped_gemm_quant.py` + `test/python/fe_api/grouped_gemm/test_grouped_gemm_quant_utils.py` (dense and discrete unified API coverage) diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_srelu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_srelu.md index d70601ac4..af78b87f0 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_srelu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_srelu.md @@ -298,5 +298,5 @@ Tuple unpacking order is: `(c_tensor, d_tensor, d_col_tensor, amax_tensor, sfd_r For end-to-end usage and regression coverage, see: -- `test/python/fe_api/test_grouped_gemm_srelu.py` -- `test/python/fe_api/test_grouped_gemm_srelu_utils.py` +- `test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu.py` +- `test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu_utils.py` diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md index 6bb8a1a34..1ac83b772 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md @@ -386,4 +386,4 @@ Returns a `TupleDict` - a dictionary-like object that also supports tuple unpack ## Usage Examples -For usage examples, see test cases in `test/python/fe_api/test_grouped_gemm_swiglu.py` + `test/python/fe_api/test_grouped_gemm_swiglu_utils.py` +For usage examples, see test cases in `test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu.py` + `test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_utils.py` diff --git a/docs/fe-oss-apis/rmsnorm_rht_amax.md b/docs/fe-oss-apis/rmsnorm_rht_amax.md index ba875fef8..caf9c5821 100644 --- a/docs/fe-oss-apis/rmsnorm_rht_amax.md +++ b/docs/fe-oss-apis/rmsnorm_rht_amax.md @@ -148,4 +148,4 @@ Tuple unpacking order is `(o_tensor, amax_tensor)`. ## Verification Focused correctness and cache coverage live in: -- `test/python/fe_api/test_rmsnorm_rht_amax.py` +- `test/python/fe_api/norm/test_rmsnorm_rht_amax.py` diff --git a/python/cudnn/native_sparse_attention/sparse_attention.md b/python/cudnn/native_sparse_attention/sparse_attention.md index 95ddd005c..87a6f487f 100644 --- a/python/cudnn/native_sparse_attention/sparse_attention.md +++ b/python/cudnn/native_sparse_attention/sparse_attention.md @@ -36,10 +36,10 @@ pip install nvidia-cudnn-frontend[cutedsl] ## Usage Sample usage and tests can be found in the (test/python) folder: -- [test_NSA_selection_attention.py](test/python/fe_api/nsa/test_NSA_selection_attention.py), `pytest test/python/test_NSA_selection_attention.py` -- [test_NSA_topk_reduction.py](test/python/fe_api/nsa/test_NSA_topk_reduction.py), `pytest test/python/test_NSA_topk_reduction.py` -- [test_NSA_compression_attention.py](test/python/fe_api/nsa/test_NSA_compression_attention.py), `pytest test/python/test_NSA_compression_attention.py` -- [test_NSA_swa.py](test/python/fe_api/nsa/test_NSA_swa.py), `pytest test/python/test_NSA_swa.py` +- [test_NSA_selection_attention.py](test/python/fe_api/nsa/test_NSA_selection_attention.py), `pytest test/python/fe_api/nsa/test_NSA_selection_attention.py` +- [test_NSA_topk_reduction.py](test/python/fe_api/nsa/test_NSA_topk_reduction.py), `pytest test/python/fe_api/nsa/test_NSA_topk_reduction.py` +- [test_NSA_compression_attention.py](test/python/fe_api/nsa/test_NSA_compression_attention.py), `pytest test/python/fe_api/nsa/test_NSA_compression_attention.py` +- [test_NSA_swa.py](test/python/fe_api/nsa/test_NSA_swa.py), `pytest test/python/fe_api/nsa/test_NSA_swa.py` Once all components are implemented, we will offer a central NSA API that will do the full NSA computation end-to-end. We will also offer the individual components as standalone APIs, as demonstrated below. diff --git a/skills/cutedsl-kernel-integration/references/integration-pattern.md b/skills/cutedsl-kernel-integration/references/integration-pattern.md index 169e47f62..608f29379 100644 --- a/skills/cutedsl-kernel-integration/references/integration-pattern.md +++ b/skills/cutedsl-kernel-integration/references/integration-pattern.md @@ -136,9 +136,9 @@ Add focused pytest coverage under `test/python/fe_api/`. Typical files: -- `test/python/fe_api/test_.py` -- `test/python/fe_api/test__utils.py` for reusable test helpers or shape/reference utilities. -- A family subdirectory when matching existing structure, such as `test/python/fe_api/nsa/`. +- `test/python/fe_api//test_.py` +- `test/python/fe_api//test__utils.py` for reusable test helpers or shape/reference utilities. +- A nested family subdirectory when matching existing structure, such as `test/python/fe_api/nsa/`. Coverage should include: diff --git a/test/python/fe_api/README.md b/test/python/fe_api/README.md new file mode 100644 index 000000000..4aa0f55d8 --- /dev/null +++ b/test/python/fe_api/README.md @@ -0,0 +1,11 @@ +# FE OSS tests + +Tests are organized by feature: + +- `gemm/`: dense GEMM fusions +- `grouped_gemm/`: grouped and discrete GEMM fusions +- `sdpa/`: scaled dot-product attention +- `bsa/`: Block Sparse Attention +- `dsa/`: DeepSeek Sparse Attention +- `nsa/`: Native Sparse Attention +- `norm/`: normalization fusions diff --git a/test/python/fe_api/block_sparse_attention/__init__.py b/test/python/fe_api/bsa/__init__.py similarity index 100% rename from test/python/fe_api/block_sparse_attention/__init__.py rename to test/python/fe_api/bsa/__init__.py diff --git a/test/python/fe_api/block_sparse_attention/bsa_reference.py b/test/python/fe_api/bsa/bsa_reference.py similarity index 100% rename from test/python/fe_api/block_sparse_attention/bsa_reference.py rename to test/python/fe_api/bsa/bsa_reference.py diff --git a/test/python/fe_api/block_sparse_attention/bsa_utils.py b/test/python/fe_api/bsa/bsa_utils.py similarity index 100% rename from test/python/fe_api/block_sparse_attention/bsa_utils.py rename to test/python/fe_api/bsa/bsa_utils.py diff --git a/test/python/fe_api/block_sparse_attention/test_BSA_attention_backward.py b/test/python/fe_api/bsa/test_BSA_attention_backward.py similarity index 97% rename from test/python/fe_api/block_sparse_attention/test_BSA_attention_backward.py rename to test/python/fe_api/bsa/test_BSA_attention_backward.py index 3ae89f0e3..703e1c561 100644 --- a/test/python/fe_api/block_sparse_attention/test_BSA_attention_backward.py +++ b/test/python/fe_api/bsa/test_BSA_attention_backward.py @@ -7,8 +7,8 @@ import torch from test_utils import torch_fork_set_rng -from fe_api.block_sparse_attention.bsa_reference import attention_backward_reference, block_sparse_mask -from fe_api.block_sparse_attention.bsa_utils import make_fixed_metadata, supported_block_size +from fe_api.bsa.bsa_reference import attention_backward_reference, block_sparse_mask +from fe_api.bsa.bsa_utils import make_fixed_metadata, supported_block_size pytestmark = [pytest.mark.gpu_exclusive, pytest.mark.xdist_group(name="gpu_exclusive")] diff --git a/test/python/fe_api/block_sparse_attention/test_BSA_attention_forward.py b/test/python/fe_api/bsa/test_BSA_attention_forward.py similarity index 98% rename from test/python/fe_api/block_sparse_attention/test_BSA_attention_forward.py rename to test/python/fe_api/bsa/test_BSA_attention_forward.py index f8ab63359..0433525d4 100644 --- a/test/python/fe_api/block_sparse_attention/test_BSA_attention_forward.py +++ b/test/python/fe_api/bsa/test_BSA_attention_forward.py @@ -8,8 +8,8 @@ import torch from test_utils import torch_fork_set_rng -from fe_api.block_sparse_attention.bsa_reference import attention_reference, block_sparse_mask -from fe_api.block_sparse_attention.bsa_utils import make_fixed_metadata, make_variable_metadata, supported_block_size +from fe_api.bsa.bsa_reference import attention_reference, block_sparse_mask +from fe_api.bsa.bsa_utils import make_fixed_metadata, make_variable_metadata, supported_block_size pytestmark = [pytest.mark.gpu_exclusive, pytest.mark.xdist_group(name="gpu_exclusive")] diff --git a/test/python/fe_api/test_gemm_amax.py b/test/python/fe_api/gemm/test_gemm_amax.py similarity index 97% rename from test/python/fe_api/test_gemm_amax.py rename to test/python/fe_api/gemm/test_gemm_amax.py index 49d71502d..58b5aad27 100644 --- a/test/python/fe_api/test_gemm_amax.py +++ b/test/python/fe_api/gemm/test_gemm_amax.py @@ -3,7 +3,7 @@ import pytest from test_utils import torch_fork_set_rng -from fe_api.test_gemm_amax_utils import ( +from fe_api.gemm.test_gemm_amax_utils import ( with_gemm_amax_params_fp4, with_gemm_amax_params_fp8, ) @@ -155,7 +155,7 @@ def _test_gemm_amax_compile_execute( try: from cudnn import GemmAmaxSm100 from cuda.bindings import driver as cuda - from fe_api.test_gemm_amax_utils import ( + from fe_api.gemm.test_gemm_amax_utils import ( allocate_input_tensors, allocate_output_tensors, check_ref_gemm_amax, @@ -242,7 +242,7 @@ def _test_gemm_amax_wrapper( try: from cudnn import gemm_amax_wrapper_sm100 from cuda.bindings import driver as cuda - from fe_api.test_gemm_amax_utils import ( + from fe_api.gemm.test_gemm_amax_utils import ( allocate_input_tensors, allocate_output_tensors, check_ref_gemm_amax, diff --git a/test/python/fe_api/test_gemm_amax_utils.py b/test/python/fe_api/gemm/test_gemm_amax_utils.py similarity index 98% rename from test/python/fe_api/test_gemm_amax_utils.py rename to test/python/fe_api/gemm/test_gemm_amax_utils.py index 9c1a8bff5..e97068805 100644 --- a/test/python/fe_api/test_gemm_amax_utils.py +++ b/test/python/fe_api/gemm/test_gemm_amax_utils.py @@ -10,7 +10,7 @@ _bfloat16_to_float4_e2m1fn_x2, float4_e2m1fn_x2_to_float32, ) -from test_fe_api_utils import create_and_permute_tensor, create_scale_factor_tensor +from fe_api.test_fe_api_utils import create_and_permute_tensor, create_scale_factor_tensor # Parameterization marks for GEMM Amax GEMM_AMAX_PARAM_MARKS_FP4 = [ diff --git a/test/python/fe_api/test_gemm_dsrelu.py b/test/python/fe_api/gemm/test_gemm_dsrelu.py similarity index 99% rename from test/python/fe_api/test_gemm_dsrelu.py rename to test/python/fe_api/gemm/test_gemm_dsrelu.py index 10f99fa00..be2e262fd 100644 --- a/test/python/fe_api/test_gemm_dsrelu.py +++ b/test/python/fe_api/gemm/test_gemm_dsrelu.py @@ -4,7 +4,7 @@ import cudnn from test_utils import torch_fork_set_rng -from fe_api.test_gemm_dsrelu_utils import ( +from fe_api.gemm.test_gemm_dsrelu_utils import ( allocate_gemm_dsrelu_outputs, allocate_gemm_dsrelu_tensors, check_ref_gemm_dsrelu, diff --git a/test/python/fe_api/test_gemm_dsrelu_utils.py b/test/python/fe_api/gemm/test_gemm_dsrelu_utils.py similarity index 99% rename from test/python/fe_api/test_gemm_dsrelu_utils.py rename to test/python/fe_api/gemm/test_gemm_dsrelu_utils.py index 764257b17..d78babfd3 100644 --- a/test/python/fe_api/test_gemm_dsrelu_utils.py +++ b/test/python/fe_api/gemm/test_gemm_dsrelu_utils.py @@ -1,7 +1,7 @@ import pytest import torch -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( compute_reference_amax, create_and_permute_tensor, create_scale_factor_tensor, diff --git a/test/python/fe_api/test_gemm_srelu.py b/test/python/fe_api/gemm/test_gemm_srelu.py similarity index 99% rename from test/python/fe_api/test_gemm_srelu.py rename to test/python/fe_api/gemm/test_gemm_srelu.py index 55ad46b83..e35cd737d 100644 --- a/test/python/fe_api/test_gemm_srelu.py +++ b/test/python/fe_api/gemm/test_gemm_srelu.py @@ -4,7 +4,7 @@ import cudnn from test_utils import torch_fork_set_rng -from fe_api.test_gemm_srelu_utils import ( +from fe_api.gemm.test_gemm_srelu_utils import ( allocate_gemm_srelu_outputs, allocate_gemm_srelu_tensors, check_ref_gemm_srelu, diff --git a/test/python/fe_api/test_gemm_srelu_utils.py b/test/python/fe_api/gemm/test_gemm_srelu_utils.py similarity index 99% rename from test/python/fe_api/test_gemm_srelu_utils.py rename to test/python/fe_api/gemm/test_gemm_srelu_utils.py index 1a47bc08b..ee4457c79 100644 --- a/test/python/fe_api/test_gemm_srelu_utils.py +++ b/test/python/fe_api/gemm/test_gemm_srelu_utils.py @@ -1,7 +1,7 @@ import pytest import torch -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( compute_reference_amax, create_and_permute_tensor, create_scale_factor_tensor, diff --git a/test/python/fe_api/test_gemm_swiglu.py b/test/python/fe_api/gemm/test_gemm_swiglu.py similarity index 99% rename from test/python/fe_api/test_gemm_swiglu.py rename to test/python/fe_api/gemm/test_gemm_swiglu.py index b41d73386..359bf2b07 100644 --- a/test/python/fe_api/test_gemm_swiglu.py +++ b/test/python/fe_api/gemm/test_gemm_swiglu.py @@ -2,7 +2,7 @@ import pytest from test_utils import torch_fork_set_rng -from fe_api.test_gemm_swiglu_utils import ( +from fe_api.gemm.test_gemm_swiglu_utils import ( allocate_input_tensors, allocate_output_tensors, check_ref_gemm_swiglu, diff --git a/test/python/fe_api/test_gemm_swiglu_utils.py b/test/python/fe_api/gemm/test_gemm_swiglu_utils.py similarity index 99% rename from test/python/fe_api/test_gemm_swiglu_utils.py rename to test/python/fe_api/gemm/test_gemm_swiglu_utils.py index 4fa8495b7..ef01bfb19 100644 --- a/test/python/fe_api/test_gemm_swiglu_utils.py +++ b/test/python/fe_api/gemm/test_gemm_swiglu_utils.py @@ -6,7 +6,7 @@ import torch import pytest from typing import Optional, Tuple -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( compute_reference_amax, create_and_permute_tensor, create_scale_factor_tensor, diff --git a/test/python/fe_api/test_discrete_grouped_gemm_dswiglu.py b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu.py similarity index 99% rename from test/python/fe_api/test_discrete_grouped_gemm_dswiglu.py rename to test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu.py index 1f44a0142..2bc1b051e 100644 --- a/test/python/fe_api/test_discrete_grouped_gemm_dswiglu.py +++ b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu.py @@ -6,7 +6,7 @@ import pytest from test_utils import torch_fork_set_rng from fe_api.test_fe_api_utils import DYNAMIC_SHAPES_M_VALUES -from fe_api.test_discrete_grouped_gemm_dswiglu_utils import ( +from fe_api.grouped_gemm.test_discrete_grouped_gemm_dswiglu_utils import ( discrete_dswiglu_init, with_discrete_dswiglu_params_fp4, with_discrete_dswiglu_params_fp8, diff --git a/test/python/fe_api/test_discrete_grouped_gemm_dswiglu_utils.py b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu_utils.py similarity index 99% rename from test/python/fe_api/test_discrete_grouped_gemm_dswiglu_utils.py rename to test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu_utils.py index d6eaa2e25..1f5edae27 100644 --- a/test/python/fe_api/test_discrete_grouped_gemm_dswiglu_utils.py +++ b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu_utils.py @@ -5,12 +5,12 @@ import torch import pytest from typing import Tuple, List, Dict, Any, Optional -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( compute_reference_amax, create_and_permute_tensor, create_scale_factor_tensor, ) -from fe_api.test_grouped_gemm_dswiglu_utils import compute_reference_row_quant +from fe_api.grouped_gemm.test_grouped_gemm_dswiglu_utils import compute_reference_row_quant # ============================================================================= # Parameterization Marks diff --git a/test/python/fe_api/test_discrete_grouped_gemm_swiglu.py b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu.py similarity index 99% rename from test/python/fe_api/test_discrete_grouped_gemm_swiglu.py rename to test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu.py index 1ebdf7880..e407095e0 100644 --- a/test/python/fe_api/test_discrete_grouped_gemm_swiglu.py +++ b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu.py @@ -9,7 +9,7 @@ import pytest from test_utils import torch_fork_set_rng from fe_api.test_fe_api_utils import DYNAMIC_SHAPES_M_VALUES -from fe_api.test_discrete_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_discrete_grouped_gemm_swiglu_utils import ( discrete_grouped_gemm_init, with_discrete_grouped_gemm_params_fp4, with_discrete_grouped_gemm_params_fp8, @@ -466,8 +466,8 @@ def test_discrete_vs_contiguous_match(ab_dtype, d_dtype, sf_vec_size, sf_dtype, if major < 10: pytest.skip(f"Requires SM100+, found SM{major}0") - from test_fe_api_utils import create_and_permute_tensor, create_scale_factor_tensor - from fe_api.test_discrete_grouped_gemm_swiglu_utils import create_mask + from fe_api.test_fe_api_utils import create_and_permute_tensor, create_scale_factor_tensor + from fe_api.grouped_gemm.test_discrete_grouped_gemm_swiglu_utils import create_mask from cudnn.api_base import ceil_div n, k, num_experts = 512, 512, 4 diff --git a/test/python/fe_api/test_discrete_grouped_gemm_swiglu_utils.py b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu_utils.py similarity index 99% rename from test/python/fe_api/test_discrete_grouped_gemm_swiglu_utils.py rename to test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu_utils.py index eb29b35ab..04efb47b5 100644 --- a/test/python/fe_api/test_discrete_grouped_gemm_swiglu_utils.py +++ b/test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu_utils.py @@ -6,7 +6,7 @@ import torch import pytest from typing import Tuple, List, Dict, Any, Optional -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( compute_reference_amax, create_and_permute_tensor, create_scale_factor_tensor, diff --git a/test/python/fe_api/test_grouped_gemm_dglu.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dglu.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_dglu.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_dglu.py index 4f3730f69..0f7757d17 100644 --- a/test/python/fe_api/test_grouped_gemm_dglu.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dglu.py @@ -9,11 +9,11 @@ import pytest from test_utils import torch_fork_set_rng from fe_api.test_fe_api_utils import DYNAMIC_SHAPES_M_VALUES -from fe_api.test_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( grouped_gemm_swiglu_init, allocate_grouped_gemm_input_tensors as allocate_grouped_gemm_input_tensors_base, ) -from fe_api.test_grouped_gemm_dswiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_dswiglu_utils import ( with_grouped_gemm_dswiglu_params_fp4, with_grouped_gemm_dswiglu_params_fp8, with_grouped_gemm_dswiglu_params_dbias_fp4, @@ -21,7 +21,7 @@ allocate_grouped_gemm_dswiglu_tensors, check_ref_grouped_gemm_dswiglu, ) -from fe_api.test_discrete_grouped_gemm_dswiglu_utils import ( +from fe_api.grouped_gemm.test_discrete_grouped_gemm_dswiglu_utils import ( discrete_dswiglu_init, allocate_discrete_dswiglu_input_tensors, allocate_discrete_dswiglu_output_tensors, diff --git a/test/python/fe_api/test_grouped_gemm_dsrelu.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_dsrelu.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu.py index e93c2d26b..02f8ed1dd 100644 --- a/test/python/fe_api/test_grouped_gemm_dsrelu.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu.py @@ -8,7 +8,7 @@ import torch import pytest from test_utils import torch_fork_set_rng -from fe_api.test_grouped_gemm_dsrelu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_dsrelu_utils import ( with_grouped_gemm_dsrelu_params_fp4, with_grouped_gemm_dsrelu_params_fp8, allocate_grouped_gemm_dsrelu_tensors, @@ -16,7 +16,7 @@ check_ref_grouped_gemm_dsrelu, grouped_gemm_dsrelu_init, ) -from fe_api.test_discrete_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_discrete_grouped_gemm_swiglu_utils import ( allocate_discrete_input_tensors, discrete_grouped_gemm_init, ) diff --git a/test/python/fe_api/test_grouped_gemm_dsrelu_utils.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu_utils.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_dsrelu_utils.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu_utils.py index 4e6c3ce59..406058191 100644 --- a/test/python/fe_api/test_grouped_gemm_dsrelu_utils.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dsrelu_utils.py @@ -6,7 +6,7 @@ import torch import pytest from typing import Optional, Tuple, List, Dict, Any -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( ceil_div, compute_reference_amax, create_and_permute_tensor, @@ -14,7 +14,7 @@ create_sf_layout_tensor, cvt_sf_MKL_to_M32x4xrm_K4xrk_L, ) -from test_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( allocate_grouped_gemm_input_tensors as allocate_grouped_gemm_input_tensors_base, get_dtype_rcp_limits as get_grouped_gemm_dtype_rcp_limits, grouped_gemm_swiglu_init as grouped_gemm_dsrelu_init, diff --git a/test/python/fe_api/test_grouped_gemm_dswiglu.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_dswiglu.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu.py index 49f0fe3fd..082a895a9 100644 --- a/test/python/fe_api/test_grouped_gemm_dswiglu.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu.py @@ -8,11 +8,11 @@ import torch import pytest from test_utils import torch_fork_set_rng -from fe_api.test_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( grouped_gemm_swiglu_init, allocate_grouped_gemm_input_tensors as allocate_grouped_gemm_input_tensors_base, ) -from fe_api.test_grouped_gemm_dswiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_dswiglu_utils import ( with_grouped_gemm_dswiglu_params_fp4, with_grouped_gemm_dswiglu_params_fp8, allocate_grouped_gemm_dswiglu_tensors, diff --git a/test/python/fe_api/test_grouped_gemm_dswiglu_utils.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_utils.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_dswiglu_utils.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_utils.py index e430af7c1..379ffd034 100644 --- a/test/python/fe_api/test_grouped_gemm_dswiglu_utils.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_utils.py @@ -6,7 +6,7 @@ import torch import pytest from typing import Optional, Tuple, List, Dict, Any -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( ceil_div, compute_reference_amax, create_and_permute_tensor, @@ -14,7 +14,7 @@ create_sf_layout_tensor, cvt_sf_MKL_to_M32x4xrm_K4xrk_L, ) -from test_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( get_dtype_rcp_limits, ) diff --git a/test/python/fe_api/test_grouped_gemm_glu.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_glu.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_glu.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_glu.py index 4c149446e..732b45131 100644 --- a/test/python/fe_api/test_grouped_gemm_glu.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_glu.py @@ -9,7 +9,7 @@ import pytest from test_utils import torch_fork_set_rng from fe_api.test_fe_api_utils import DYNAMIC_SHAPES_M_VALUES -from fe_api.test_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( grouped_gemm_swiglu_init, with_grouped_gemm_swiglu_params_fp4, with_grouped_gemm_swiglu_params_fp8, @@ -19,7 +19,7 @@ allocate_grouped_gemm_output_tensors, check_ref_grouped_gemm_swiglu, ) -from fe_api.test_discrete_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_discrete_grouped_gemm_swiglu_utils import ( discrete_grouped_gemm_init, allocate_discrete_input_tensors, allocate_discrete_output_tensors, diff --git a/test/python/fe_api/test_grouped_gemm_glu_hadamard.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_glu_hadamard.py similarity index 98% rename from test/python/fe_api/test_grouped_gemm_glu_hadamard.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_glu_hadamard.py index 48c847986..932c14cd5 100644 --- a/test/python/fe_api/test_grouped_gemm_glu_hadamard.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_glu_hadamard.py @@ -6,9 +6,9 @@ import torch from test_utils import torch_fork_set_rng -from fe_api.test_discrete_grouped_gemm_swiglu_utils import allocate_discrete_input_tensors +from fe_api.grouped_gemm.test_discrete_grouped_gemm_swiglu_utils import allocate_discrete_input_tensors from fe_api.test_fe_api_utils import DYNAMIC_SHAPES_M_VALUES, compute_reference_amax -from fe_api.test_grouped_gemm_swiglu_utils import allocate_grouped_gemm_input_tensors, grouped_gemm_swiglu_init +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import allocate_grouped_gemm_input_tensors, grouped_gemm_swiglu_init FP4_EXECUTION_CASES = [ (torch.float4_e2m1fn_x2, torch.float8_e8m0fnu, 16), diff --git a/test/python/fe_api/test_grouped_gemm_quant.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_quant.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_quant.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_quant.py index 88b9092ab..e8ec2f7d2 100644 --- a/test/python/fe_api/test_grouped_gemm_quant.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_quant.py @@ -10,13 +10,13 @@ import pytest from test_utils import torch_fork_set_rng from fe_api.test_fe_api_utils import DYNAMIC_SHAPES_M_VALUES -from fe_api.test_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( allocate_grouped_gemm_input_tensors, ) -from fe_api.test_discrete_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_discrete_grouped_gemm_swiglu_utils import ( allocate_discrete_input_tensors, ) -from fe_api.test_grouped_gemm_quant_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_quant_utils import ( grouped_gemm_quant_init, with_grouped_gemm_quant_params_fp4, with_grouped_gemm_quant_params_fp8, diff --git a/test/python/fe_api/test_grouped_gemm_quant_utils.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_quant_utils.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_quant_utils.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_quant_utils.py index 080c442b2..a1c9160e7 100644 --- a/test/python/fe_api/test_grouped_gemm_quant_utils.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_quant_utils.py @@ -8,7 +8,7 @@ import torch import pytest from typing import Optional, Tuple, List, Dict, Any -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( ceil_div, compute_reference_amax, create_and_permute_tensor, diff --git a/test/python/fe_api/test_grouped_gemm_srelu.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_srelu.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu.py index 8061a53f9..402803e84 100644 --- a/test/python/fe_api/test_grouped_gemm_srelu.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu.py @@ -10,7 +10,7 @@ import torch import pytest from test_utils import torch_fork_set_rng -from fe_api.test_grouped_gemm_srelu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_srelu_utils import ( grouped_gemm_srelu_init, with_grouped_gemm_srelu_params_fp4, with_grouped_gemm_srelu_params_fp8, @@ -18,7 +18,7 @@ allocate_grouped_gemm_output_tensors, check_ref_grouped_gemm_srelu, ) -from fe_api.test_discrete_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_discrete_grouped_gemm_swiglu_utils import ( allocate_discrete_input_tensors, discrete_grouped_gemm_init, ) diff --git a/test/python/fe_api/test_grouped_gemm_srelu_utils.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu_utils.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_srelu_utils.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu_utils.py index 63a24d29c..c1c059a87 100644 --- a/test/python/fe_api/test_grouped_gemm_srelu_utils.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu_utils.py @@ -8,7 +8,7 @@ import torch import pytest from typing import Optional, Tuple, List, Dict, Any -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( ceil_div, compute_reference_amax, create_and_permute_tensor, diff --git a/test/python/fe_api/test_grouped_gemm_swiglu.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_swiglu.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu.py index 7eee14a85..28d262b6a 100644 --- a/test/python/fe_api/test_grouped_gemm_swiglu.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu.py @@ -10,7 +10,7 @@ import torch import pytest from test_utils import torch_fork_set_rng -from fe_api.test_grouped_gemm_swiglu_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_swiglu_utils import ( grouped_gemm_swiglu_init, with_grouped_gemm_swiglu_params_fp4, with_grouped_gemm_swiglu_params_fp8, diff --git a/test/python/fe_api/test_grouped_gemm_swiglu_utils.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_utils.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_swiglu_utils.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_utils.py index 28757a121..e81eb919f 100644 --- a/test/python/fe_api/test_grouped_gemm_swiglu_utils.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_utils.py @@ -8,7 +8,7 @@ import torch import pytest from typing import Optional, Tuple, List, Dict, Any -from test_fe_api_utils import ( +from fe_api.test_fe_api_utils import ( ceil_div, compute_reference_amax, create_and_permute_tensor, diff --git a/test/python/fe_api/test_grouped_gemm_wgrad.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_wgrad.py similarity index 99% rename from test/python/fe_api/test_grouped_gemm_wgrad.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_wgrad.py index 6c8b7765a..c8952f1fc 100644 --- a/test/python/fe_api/test_grouped_gemm_wgrad.py +++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_wgrad.py @@ -5,7 +5,7 @@ import cudnn from test_utils import torch_fork_set_rng -from fe_api.test_grouped_gemm_wgrad_utils import ( +from fe_api.grouped_gemm.test_grouped_gemm_wgrad_utils import ( grouped_gemm_wgrad_init, with_grouped_gemm_wgrad_params_fp4, with_grouped_gemm_wgrad_params_fp8, diff --git a/test/python/fe_api/test_grouped_gemm_wgrad_utils.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_wgrad_utils.py similarity index 100% rename from test/python/fe_api/test_grouped_gemm_wgrad_utils.py rename to test/python/fe_api/grouped_gemm/test_grouped_gemm_wgrad_utils.py diff --git a/test/python/fe_api/test_rmsnorm_rht_amax.py b/test/python/fe_api/norm/test_rmsnorm_rht_amax.py similarity index 100% rename from test/python/fe_api/test_rmsnorm_rht_amax.py rename to test/python/fe_api/norm/test_rmsnorm_rht_amax.py diff --git a/test/python/fe_api/test_sdpa_bwd.py b/test/python/fe_api/sdpa/test_sdpa_bwd.py similarity index 99% rename from test/python/fe_api/test_sdpa_bwd.py rename to test/python/fe_api/sdpa/test_sdpa_bwd.py index 88f8950a1..26800ca2b 100644 --- a/test/python/fe_api/test_sdpa_bwd.py +++ b/test/python/fe_api/sdpa/test_sdpa_bwd.py @@ -11,7 +11,7 @@ import torch from test_utils import torch_fork_set_rng -from fe_api.test_sdpa_bwd_utils import ( +from fe_api.sdpa.test_sdpa_bwd_utils import ( allocate_sdpa_bwd_input_tensors, allocate_sdpa_bwd_output_tensors, check_ref_sdpa_bwd, diff --git a/test/python/fe_api/test_sdpa_bwd_utils.py b/test/python/fe_api/sdpa/test_sdpa_bwd_utils.py similarity index 100% rename from test/python/fe_api/test_sdpa_bwd_utils.py rename to test/python/fe_api/sdpa/test_sdpa_bwd_utils.py diff --git a/test/python/fe_api/test_sdpa_fwd.py b/test/python/fe_api/sdpa/test_sdpa_fwd.py similarity index 98% rename from test/python/fe_api/test_sdpa_fwd.py rename to test/python/fe_api/sdpa/test_sdpa_fwd.py index 32f8aa556..339fd58ce 100644 --- a/test/python/fe_api/test_sdpa_fwd.py +++ b/test/python/fe_api/sdpa/test_sdpa_fwd.py @@ -6,7 +6,7 @@ import torch from test_utils import torch_fork_set_rng -from fe_api.test_sdpa_fwd_utils import ( +from fe_api.sdpa.test_sdpa_fwd_utils import ( allocate_sdpa_fwd_input_tensors, allocate_sdpa_fwd_output_tensors, check_ref_sdpa_fwd, diff --git a/test/python/fe_api/test_sdpa_fwd_utils.py b/test/python/fe_api/sdpa/test_sdpa_fwd_utils.py similarity index 100% rename from test/python/fe_api/test_sdpa_fwd_utils.py rename to test/python/fe_api/sdpa/test_sdpa_fwd_utils.py