From efedd04da85800cf4cee94ec3462fb2126ebe20b Mon Sep 17 00:00:00 2001 From: Christophe Guillon Date: Thu, 17 Sep 2026 17:10:26 +0200 Subject: [PATCH] db: update json db to format 0.3 with backend kwargs --- src/xtc/cli/query_results.py | 17 ++++++++++++++--- src/xtc/search/callback.py | 7 ++++++- src/xtc/search/explore.py | 14 ++++++++++---- 3 files changed, 30 insertions(+), 8 deletions(-) diff --git a/src/xtc/cli/query_results.py b/src/xtc/cli/query_results.py index 57c7149e..e4190905 100644 --- a/src/xtc/cli/query_results.py +++ b/src/xtc/cli/query_results.py @@ -23,7 +23,7 @@ class ResultsDB(ABC): - VERSION = "v0.2" + VERSION = "v0.3" def __init__( self, @@ -49,6 +49,7 @@ def _default_match( target: str = "native", threads: int = 1, backend: str | None = None, + backend_kwargs: dict[str, Any] | None = None, ) -> Generator[DBEntry, None, None]: version = self.get_version() if self._node_target == "native": @@ -58,7 +59,7 @@ def _default_match( f"node must be specified for non native target" ) platform = self.get_node_platform(self._node, self._node_target) - compiler = self.get_compiler(target, threads, backend) + compiler = self.get_compiler(target, threads, backend, backend_kwargs) operator = self.get_operator(graph) logger.debug("MATCH: version: %s", version) logger.debug("MATCH: platform: %s", platform) @@ -84,11 +85,17 @@ def get_version(cls) -> list[Any]: @classmethod def get_compiler( - cls, target: str = "native", threads: int = 1, backend: str | None = None + cls, + target: str = "native", + threads: int = 1, + backend: str | None = None, + backend_kwargs: dict[str, Any] | None = None, ) -> list[Any]: compiler = ["xtc", cls.get_xtc_version(), target, threads] if backend is not None: compiler.append(backend) + if backend_kwargs is not None: + compiler.append(backend_kwargs) return compiler @classmethod @@ -128,6 +135,7 @@ def get_results( target: str = "native", threads: int = 1, backend: str | None = None, + backend_kwargs: dict[str, Any] | None = None, allow_errors: bool = False, ) -> list[DBEntry]: results = [] @@ -137,6 +145,7 @@ def get_results( target=target, threads=threads, backend=backend, + backend_kwargs=backend_kwargs, ): if not allow_errors and log["results"][0] != 0: continue @@ -149,12 +158,14 @@ def get_best( target: str = "native", threads: int = 1, backend: str | None = None, + backend_kwargs: dict[str, Any] | None = None, ) -> DBEntry | None: logs = self.get_results( graph, target, threads, backend, + backend_kwargs, allow_errors=False, ) ordered = sorted(logs, key=lambda x: min(x["results"][1])) diff --git a/src/xtc/search/callback.py b/src/xtc/search/callback.py index ad433a4f..b5d08053 100644 --- a/src/xtc/search/callback.py +++ b/src/xtc/search/callback.py @@ -31,6 +31,7 @@ def __init__( target: str, threads: int, strategy: str, + backend_kwargs: dict[str, dict[str, Any]] | None = None, ) -> None: self._dbfile = dbfile self._target = target @@ -39,6 +40,7 @@ def __init__( self._platform = ResultsDB.get_native_platform() self._operator: list[Any] | None = None self._strategy = ResultsDB.get_strategy(strategy) + self._backend_kwargs = backend_kwargs if backend_kwargs is not None else {} def set_graph(self, graph: Graph): # assert len(graph.nodes) == 1, f"Only support recording of single node graph" @@ -49,7 +51,10 @@ def _write_result(self, result: Sequence) -> None: x, code, time, backend = result if code != 0: time = 0 - compiler = ResultsDB.get_compiler(self._target, self._threads, backend) + backend_kwargs = self._backend_kwargs.get(backend, {}) + compiler = ResultsDB.get_compiler( + self._target, self._threads, backend, backend_kwargs + ) log = dict( version=self._version, platform=self._platform, diff --git a/src/xtc/search/explore.py b/src/xtc/search/explore.py index 14f09cc3..894df9a1 100644 --- a/src/xtc/search/explore.py +++ b/src/xtc/search/explore.py @@ -113,6 +113,7 @@ class ExplorationConfig: use_tensors: bool = False progress_cls: str = "tqdm" module_type: str = "shlib" + backend_kwargs: dict[str, dict[str, Any]] = field(default_factory=dict) def __post_init__(self): if self.graph_file is not None: @@ -133,6 +134,13 @@ def __post_init__(self): f"backend {backend} not available for operator {self.operator}" ) + for backend in self.backends: + if backend not in self.backend_kwargs: + self.backend_kwargs[backend] = {} + + if self.use_tensors: + self.backend_kwargs["mlir"].update({"use_tensor_dialect": True}) + @staticmethod def from_args( args: NS | None = None, @@ -352,13 +360,10 @@ def compile_one( args = self.config assert isinstance(in_x, list), f"X not a list: {in_x} ({type(in_x)})" logger.debug("Compile: %s: %s: %s...", ident, backend, in_x) - kwargs = {} - if backend == "mlir": - kwargs.update({"use_tensor_dialect": args.use_tensors}) impl, backend_name = self.graph_implementer( graph, backend, - **kwargs, + **args.backend_kwargs.get(backend, {}), ) assert backend_name == backend scheduler = impl.get_scheduler() @@ -714,6 +719,7 @@ def get_result_callbacks( "native", args.threads, self.get_strategy_name(args.strategy), + args.backend_kwargs, ) args.db_callback.set_graph(graph) callbacks.append(args.db_callback)