diff --git a/src/mldebug/layer_info.py b/src/mldebug/layer_info.py index 12ff9fc..7135189 100644 --- a/src/mldebug/layer_info.py +++ b/src/mldebug/layer_info.py @@ -13,6 +13,7 @@ from pathlib import Path from mldebug.aie_overlay import Overlay +from mldebug.arch.device_configs import AIE_DEV_PHX, AIE_DEV_STX from mldebug.mladf_report import MladfReport from mldebug.work_dir import _parse_flexml_layer_id from mldebug.work_dir import WorkDir @@ -26,17 +27,38 @@ "mllib_graphs::resize_adf_wrapper", # This has many sublayers and needs to be better understood "mllib_graphs::mha_type1::mha_adf_wrapper", - # Causes failure. TODO: investigate - "superkernel_eltunary", # Padding preamble; halting on it desyncs PC/iteration stepping on HW "buffer_pad_innermost", - "superkernel_conv_eltbinary", # TODO: investigate why this is causing a failure (ref: AIESW-43867) "mllib_graphs::topk_adf_wrapper", "mllib_graphs::transpose4d_adf_wrapper", "mllib_graphs::transpose4d_adf_wrapper", ] +# Kernels that work on PHX/STX but are unsupported on other devices. +# Causes failure on other devices. TODO: investigate +phx_stx_supported_kernels = [ + "superkernel_eltunary", + "superkernel_conv_eltbinary", +] + +# Set once by the CLI before any LayerInfo is built +_device = None + + +def set_device(device): + """Record the base device used to filter device-specific unsupported kernels.""" + global _device # pylint: disable=global-statement + _device = device + + +def _is_unsupported_kernel(name): + """True if kernel name matches an unsupported superkernel for the current device.""" + kernels = unsupported_superkernels + if _device not in (AIE_DEV_PHX, AIE_DEV_STX): + kernels = kernels + phx_stx_supported_kernels + return any(k in name for k in kernels) + def _strip_template(name): """Strip C++ template parameters for compiler-agnostic name comparison. @@ -321,7 +343,7 @@ def __init__( # 1. Layers without any kernel should be skipped # 2. Unsupported superkernel should be skipped - if info.get("is_concat") or not kname or any(k in kname for k in unsupported_superkernels): + if info.get("is_concat") or not kname or _is_unsupported_kernel(kname): LOGGER.verbose_print( f"[WARNING] unsupported kernel {kname} at Layer {self.layer_order} will be skipped." ) @@ -336,7 +358,7 @@ def __init__( if ( not stamp.name or stamp.elf_name == -1 - or any(k in stamp.name for k in unsupported_superkernels) + or _is_unsupported_kernel(stamp.name) ): LOGGER.verbose_print( f"[WARNING] unsupported kernel {stamp.name} at Layer {self.layer_order} will be skipped." diff --git a/src/mldebug/mldebug_cli.py b/src/mldebug/mldebug_cli.py index 8ffaeea..c256c0e 100644 --- a/src/mldebug/mldebug_cli.py +++ b/src/mldebug/mldebug_cli.py @@ -23,6 +23,7 @@ AIE_DEV_TEL, ) from mldebug.client_debug import ClientDebug +from mldebug import layer_info from mldebug.input_parser import ( check_hw_context, check_registry_keys, @@ -44,8 +45,6 @@ def _apply_unsupported_kernels_from_args(args): This must happen before LayerInfo creates Layer objects (ClientDebug -> LayerInfo). """ - from mldebug import layer_info # pylint: disable=import-outside-toplevel - values = args.unsupported_kernels if not values: return @@ -190,6 +189,7 @@ def launch_debug(args, output_dir): if args.backend == "xrt": context_id, pid = check_hw_context(args) # Top debug handle + layer_info.set_device(args.device) _apply_unsupported_kernels_from_args(args) _apply_unsupported_layers_from_args(args) handle = ClientDebug(args, context_id, pid, output_dir)