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
Original file line number Diff line number Diff line change
Expand Up @@ -140,7 +140,7 @@ def _generate_lane_bodies(
[
f"reward_value_{index} = {cls._term_call(term, index, 'ctx')}",
f"reward_terms[env_id, {index}] = reward_value_{index}",
f"weighted_reward_{index} = reward_value_{index} * reward_weights[{index}]",
f"weighted_reward_{index} = reward_value_{index} * reward_weights[{index}] * ctx.dt",
f"weighted_reward_terms[env_id, {index}] = weighted_reward_{index}",
f"total_reward += weighted_reward_{index}",
]
Expand Down
185 changes: 163 additions & 22 deletions motrix_env_core/src/motrix_env_core/numba/manager/compiler/compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from pathlib import Path
from time import perf_counter
from types import MappingProxyType
from typing import Any, get_args, get_origin, get_type_hints

Expand All @@ -27,6 +28,7 @@
flatten_kernel_data,
is_kernel_data,
iter_layout_leaves,
kernel_data,
lane_expression,
map_proxy,
proxy_symbol,
Expand Down Expand Up @@ -79,6 +81,20 @@ def _invalidate_term_cache() -> None:
_TERM_CACHE.clear()


@kernel_data
class TermScalarBuffer:
"""Runtime buffer slot for one numeric term argument.

Numeric tuning arguments (thresholds, scales, gains, ...) are lowered into
a fixed-dtype shared input slot: only the type/layout fingerprint enters
the plan key, and the value itself is read at runtime. This keeps
tuning-variant manager configs on one compiled plan instead of triggering
a full recompile per value.
"""

value: np.float32


@dataclass(frozen=True)
class ManagerSimInput:
"""One compiled simulator input exposed through ``ManagerContext.sim``."""
Expand Down Expand Up @@ -111,6 +127,7 @@ def __init__(self, env: ManagerEnv):
self._plan_parts.append(("sim_inputs", sim_input_fingerprint))

def build(self) -> CompiledManagerProgram:
build_started = perf_counter()
self._resolve_runtime_terms()
observation_groups = self._env.observation_groups
reward_terms = self._env._reward_terms
Expand All @@ -130,9 +147,17 @@ def build(self) -> CompiledManagerProgram:
)
self._validate_observation_spaces(observation_groups)

kernel_reward_weights = np.asarray(reward_weights, dtype=np.float32) * np.float32(self._env.cfg.ctrl_dt)
self._plan_parts.append(("manager_context_dt", np.float32(self._env.cfg.ctrl_dt)))
# Reward weights stay unscaled in the runtime buffer; the dt scaling
# happens inside the kernel via ctx.dt, so ctrl_dt does not need to
# participate in the plan key.
kernel_reward_weights = np.asarray(reward_weights, dtype=np.float32)
plan_key = hashlib.sha256(repr(tuple(self._plan_parts)).encode()).hexdigest()
logger.info(
"Manager startup %s: build plan finished in %.3fs (plan_key=%.12s)",
self._env_name(),
perf_counter() - build_started,
plan_key,
)
source = KernelSourceGenerator(self._flat_input_count).generate(
observation_terms,
observation_layout,
Expand All @@ -146,8 +171,28 @@ def build(self) -> CompiledManagerProgram:
)
generated_filename = self._materialize_source(source, plan_key)
kernels = _KERNEL_CACHE.get(plan_key)
if kernels is None:
logger.debug("Manager kernel cache miss: plan_key=%s", plan_key)
if kernels is not None:
logger.info(
"Manager startup %s: kernels reused (cache=in-process, plan_key=%.12s)",
self._env_name(),
plan_key,
)
else:
disk_cache = self._has_generated_kernel_cache(plan_key)
if disk_cache:
logger.info(
"Manager startup %s: compiled kernel cache found, loading (cache=disk, plan_key=%.12s)",
self._env_name(),
plan_key,
)
else:
logger.info(
"Manager startup %s: no cache for this plan, first compile may take a while "
"(cache=miss, plan_key=%.12s)",
self._env_name(),
plan_key,
)
compile_started = perf_counter()
try:
kernels = self._compile_kernels(source, generated_filename)
except (ImportError, OSError, RuntimeError, TypeError, ValueError) as error:
Expand All @@ -156,8 +201,12 @@ def build(self) -> CompiledManagerProgram:
generated_filename = self._materialize_source(source, plan_key)
kernels = self._compile_kernels(source, generated_filename)
_KERNEL_CACHE[plan_key] = kernels
else:
logger.debug("Manager kernel in-process cache hit: plan_key=%s", plan_key)
logger.info(
"Manager startup %s: compile kernels finished in %.3fs (cache=%s)",
self._env_name(),
perf_counter() - compile_started,
"disk" if disk_cache else "miss",
)
evaluate_kernel, observe_kernel, reset_kernel = kernels

layout = ManagerLayout(
Expand Down Expand Up @@ -271,6 +320,7 @@ def _resolve_reset(
symbol = f"term_reset_{term_index}"
self._term_functions[symbol] = self._generated_term(function, dispatcher)
expressions = []
plan_expressions = []
prepared_indices = []
for index, value in enumerate(term.args):
if is_kernel_data(value):
Expand All @@ -285,15 +335,20 @@ def _resolve_reset(
)
prepared_indices.append(prepared_index)
expressions.append(expression)
plan_expressions.append(expression)
else:
prepared_indices.append(None)
expressions.append(repr(value))
expression, _, plan_part, prepared_index = self._resolve_plain_arg(
f"sim_reset.{name}.args[{index}]", value
)
prepared_indices.append(prepared_index)
expressions.append(expression)
plan_expressions.append(plan_part)
self._plan_parts.append(
(
"sim_reset",
name,
self._function_fingerprint(function),
tuple(expressions),
plan_expressions,
repr(descriptors),
fields,
)
Expand Down Expand Up @@ -473,7 +528,9 @@ def _resolve_observation_dispatch(
f"ManagerContext and np.ndarray (conventionally named ctx and out)."
)
expressions = []
plan_expressions = []
prepared_indices = []
warmup_values = []
for index, value in enumerate(term.args):
if is_kernel_data(value):
_, tree_def = flatten_kernel_data(value)
Expand All @@ -485,23 +542,30 @@ def _resolve_observation_dispatch(
)
prepared_indices.append(prepared_index)
expressions.append(expression)
plan_expressions.append(expression)
warmup_values.append(value)
else:
prepared_indices.append(None)
expressions.append(repr(value))
expression, warmup, plan_part, prepared_index = self._resolve_plain_arg(
f"observation.{group}.{name}.args[{index}]", value
)
prepared_indices.append(prepared_index)
expressions.append(expression)
plan_expressions.append(plan_part)
warmup_values.append(warmup)
expressions = tuple(expressions)
dispatcher = self._compile_term(function)
symbol = f"term_observation_{term_index}"
self._term_functions[symbol] = self._generated_term(function, dispatcher)
self._plan_parts.append(
(group, name, "observation_dispatch", self._function_fingerprint(function), output_size, expressions)
(group, name, "observation_dispatch", self._function_fingerprint(function), output_size, plan_expressions)
)
return PreparedInvocation(
f"{group}.{name}",
"observation",
dispatcher,
None,
args_expressions=expressions,
args_values=term.args,
args_values=tuple(warmup_values),
args_prepared_indices=tuple(prepared_indices),
output_size=output_size,
)
Expand Down Expand Up @@ -574,7 +638,9 @@ def _resolve_dispatch_term(
f"(conventionally named ctx)."
)
expressions = []
plan_expressions = []
prepared_indices = []
warmup_values = []
for index, value in enumerate(term.args):
if is_kernel_data(value):
_, tree_def = flatten_kernel_data(value)
Expand All @@ -586,20 +652,27 @@ def _resolve_dispatch_term(
)
prepared_indices.append(prepared_index)
expressions.append(expression)
plan_expressions.append(expression)
warmup_values.append(value)
else:
prepared_indices.append(None)
expressions.append(repr(value))
expression, warmup, plan_part, prepared_index = self._resolve_plain_arg(
f"{group}.{name}.args[{index}]", value
)
prepared_indices.append(prepared_index)
expressions.append(expression)
plan_expressions.append(plan_part)
warmup_values.append(warmup)
dispatcher = self._compile_term(function)
symbol = f"term_{group}_{term_index}"
self._term_functions[symbol] = self._generated_term(function, dispatcher)
self._plan_parts.append((group, name, "dispatch", self._function_fingerprint(function), expressions))
self._plan_parts.append((group, name, "dispatch", self._function_fingerprint(function), plan_expressions))
return PreparedInvocation(
f"{group}.{name}",
group,
dispatcher,
None,
args_expressions=tuple(expressions),
args_values=term.args,
args_values=tuple(warmup_values),
args_prepared_indices=tuple(prepared_indices),
)

Expand Down Expand Up @@ -681,6 +754,38 @@ def _register_prepared(
self._prepared_types[proxy_symbol(proxy)] = proxy
return prepared_index, lane_expression(layout, input_offset)

def _scalar_buffer_arg(self, source_name: str, value: float) -> tuple[str, np.float32]:
"""Lower one numeric term argument into a runtime scalar buffer slot.

Returns the kernel-lane argument expression and the normalized warmup
value. Only the fixed float32 layout enters the plan; the value itself
is read from the input slot at runtime.
"""
scalar = np.float32(value)
wrapped = TermScalarBuffer(scalar)
_, tree_def = flatten_kernel_data(wrapped)
layout = self._kernel_data_lowering.lower(
tree_def,
context=source_name,
force_shared=True,
)
prepared_index, _ = self._register_prepared(source_name, wrapped, TermScalarBuffer, layout)
leaf = iter_layout_leaves(layout)[0]
return f"input_{self._input_offsets[prepared_index] + leaf.slot_index}", scalar

def _resolve_plain_arg(self, source_name: str, value: Any) -> tuple[str | None, Any, str, int | None]:
"""Resolve one non-kernel-data term argument.

Float-like scalars become runtime scalar buffers; every
other scalar stays a compile-time constant. Returns the kernel-lane
argument expression, the warmup value, the plan fingerprint part, and
the prepared index (``None`` for both plain scalars and buffers).
"""
if isinstance(value, (float, np.floating)):
expression, warmup = self._scalar_buffer_arg(source_name, value)
return expression, warmup, "scalar_buffer(np.float32)", None
return repr(value), value, repr(value), None

def _validate_observation_spaces(self, groups: dict[str, ObservationGroupEntry]) -> None:
expected_policy = self._env.policy_observation_space.shape
if expected_policy != (groups["policy"].size,):
Expand All @@ -704,9 +809,7 @@ def _validate_observation_spaces(self, groups: dict[str, ObservationGroupEntry])

@staticmethod
def _materialize_source(source: str, plan_key: str) -> str:
cache_dir = os.environ.get("NUMBA_CACHE_DIR")
if not cache_dir:
cache_dir = os.path.join(os.path.expanduser("~"), ".cache", "motrixlab", "numba-manager")
cache_dir = NumbaKernelCompiler._generated_cache_dir()
path = Path(cache_dir) / "generated" / f"{plan_key}.py"
path.parent.mkdir(parents=True, exist_ok=True)
if not path.exists() or path.read_text(encoding="utf-8") != source:
Expand All @@ -724,23 +827,53 @@ def _materialize_source(source: str, plan_key: str) -> str:
return str(path)

@staticmethod
def _invalidate_generated_cache(plan_key: str) -> None:
def _generated_cache_dir() -> str:
cache_dir = os.environ.get("NUMBA_CACHE_DIR")
if not cache_dir:
cache_dir = os.path.join(os.path.expanduser("~"), ".cache", "motrixlab", "numba-manager")
return cache_dir

@staticmethod
def _generated_cache_paths(plan_key: str) -> tuple[Path, ...]:
"""Return compiled-kernel cache files for one plan key.

Numba stores the cached specializations of the generated module next to
its cache locator: under ``generated_<hash>/`` when ``NUMBA_CACHE_DIR``
is set, and under ``generated/__pycache__/`` (beside the generated
source) when it is not. Both layouts are covered so probe and
invalidation agree with where the files actually live.
"""
cache_dir = Path(NumbaKernelCompiler._generated_cache_dir())
paths: list[Path] = []
for pattern in (f"generated_*/*{plan_key}*", f"generated/__pycache__/*{plan_key}*"):
try:
paths.extend(cache_dir.glob(pattern))
except OSError:
logger.debug("Unable to scan generated kernel cache: %s", cache_dir, exc_info=True)
return tuple(paths)

@staticmethod
def _has_generated_kernel_cache(plan_key: str) -> bool:
"""Return whether compiled-kernel cache files exist for one plan key."""
return any(path.suffix in {".nbi", ".nbc"} for path in NumbaKernelCompiler._generated_cache_paths(plan_key))

@staticmethod
def _invalidate_generated_cache(plan_key: str) -> None:
cache_dir = NumbaKernelCompiler._generated_cache_dir()
source_path = Path(cache_dir) / "generated" / f"{plan_key}.py"
try:
source_path.unlink(missing_ok=True)
except OSError:
logger.debug("Unable to remove invalid generated source: %s", source_path, exc_info=True)
for cache_path in Path(cache_dir).glob(f"generated_*/*{plan_key}*"):
for cache_path in NumbaKernelCompiler._generated_cache_paths(plan_key):
try:
cache_path.unlink(missing_ok=True)
except OSError:
logger.debug("Unable to remove invalid generated cache: %s", cache_path, exc_info=True)

def compile_specializations(self, inputs: tuple[Any, ...]) -> None:
"""Compile all term and fused-kernel specializations for the environment."""
started = perf_counter()
try:
self._compile_specializations_once(inputs)
except (
Expand All @@ -760,6 +893,14 @@ def compile_specializations(self, inputs: tuple[Any, ...]) -> None:
_KERNEL_CACHE.pop(plan_key, None)
self._env._task_program = self.build().task
self._compile_specializations_once(inputs)
logger.info(
"Manager startup %s: term specializations finished in %.3fs",
self._env_name(),
perf_counter() - started,
)

def _env_name(self) -> str:
return type(self._env).__name__

def _compile_specializations_once(self, inputs: tuple[Any, ...]) -> None:
program = self._env._compiled_manager_program
Expand Down
Loading
Loading