"""TestSpec adapter for SparseFlashMla TTK assets."""
import importlib.util
import sys
from pathlib import Path
ASSET_IMPL_DIR = Path(__file__).with_name("impl")
def load_impl_module(stem):
name = f"smla_ttk_{stem}"
if name in sys.modules:
return sys.modules[name]
path = ASSET_IMPL_DIR / f"{stem}.py"
try:
spec = importlib.util.spec_from_file_location(name, path)
if spec is None or spec.loader is None:
raise ImportError(f"cannot create import spec for {path}")
module = importlib.util.module_from_spec(spec)
sys.modules[name] = module
spec.loader.exec_module(module)
except Exception as exc:
sys.modules.pop(name, None)
raise RuntimeError(
"Failed to load SparseFlashMla assets module; "
f"stage=impl/{stem}; module={path.resolve()}; "
f"original error: {type(exc).__name__}: {exc}"
) from exc
return module
npu_preprocess_module = load_impl_module("npu_preprocess")
golden_module = load_impl_module("golden")
inputs_module = load_impl_module("inputs")
metadata_inputs_module = load_impl_module("metadata_inputs")
compare_module = load_impl_module("compare")
class SparseFlashMlaSpec:
golden = golden_module.cpu_sparse_flash_mla
customize_inputs = inputs_module.generate_sparse_flash_mla_inputs
npu_preprocess = npu_preprocess_module.run
tolerance = {
"float16": {"standard": "stat_rel_err"},
"bfloat16": {"standard": "stat_rel_err"},
}
compare = staticmethod(compare_module.compare)
class AclnnSparseFlashMlaSpec(SparseFlashMlaSpec):
golden = golden_module.cpu_aclnn_sparse_flash_mla
customize_inputs = inputs_module.generate_aclnn_sparse_flash_mla_inputs
npu_preprocess = npu_preprocess_module.run_aclnn
class SparseFlashMlaMetadataSpec:
customize_inputs = metadata_inputs_module.generate_sparse_flash_mla_metadata_inputs
__spec__ = {
"torch.ops.cann_ops_transformer.sparse_flash_mla": "SparseFlashMlaSpec",
"torch.ops.cann_ops_transformer.sparse_flash_mla_metadata": (
"SparseFlashMlaMetadataSpec"
),
"aclnnSparseFlashMla": "AclnnSparseFlashMlaSpec",
}