Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 14 additions & 3 deletions src/xtc/cli/query_results.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@


class ResultsDB(ABC):
VERSION = "v0.2"
VERSION = "v0.3"

def __init__(
self,
Expand All @@ -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":
Expand All @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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 = []
Expand All @@ -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
Expand All @@ -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]))
Expand Down
7 changes: 6 additions & 1 deletion src/xtc/search/callback.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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"
Expand All @@ -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,
Expand Down
14 changes: 10 additions & 4 deletions src/xtc/search/explore.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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,
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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)
Expand Down
Loading