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
4 changes: 2 additions & 2 deletions python/simpler/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -619,7 +619,7 @@ def _run_chip_main_loop( # noqa: PLR0912 -- TASK_READY + 6 control sub-commands
# when a prior _CTRL_UNREGISTER failed before reaching
# prepared.discard, while the parent still popped its
# registry under best-effort semantics. Without this,
# register_prepared_callable would fail-fast on a slot the
# register_callable would fail-fast on a slot the
# user was told is reusable. The `cid in prepared` gate
# keeps the happy path at zero added cost.
if int(cid) in prepared:
Expand Down Expand Up @@ -1062,7 +1062,7 @@ def _allocate_cid(self) -> int:
# The AICPU side keeps a fixed-size orch_so_table_ keyed by cid;
# raise here so the failure surfaces at register-time with a
# protocol-aware message, not later from
# DeviceRunner::register_prepared_callable with a generic
# DeviceRunner::register_callable with a generic
# "out of range" log.
raise RuntimeError(
"Worker.register: cid space exhausted "
Expand Down
79 changes: 36 additions & 43 deletions src/a2a3/platform/onboard/host/device_runner.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -607,8 +607,8 @@ int DeviceRunner::prepare_orch_so(Runtime &runtime) {
LOG_ERROR("prepare_orch_so: no active callable_id; prepared-callable flow required");
return -1;
}
auto it = prepared_callables_.find(cid);
if (it == prepared_callables_.end()) {
auto it = callables_.find(cid);
if (it == callables_.end()) {
LOG_ERROR("prepare_orch_so: callable_id=%d not registered", cid);
return -1;
}
Expand Down Expand Up @@ -636,7 +636,7 @@ int DeviceRunner::prepare_orch_so(Runtime &runtime) {
return 0;
}

int DeviceRunner::register_prepared_callable(
int DeviceRunner::register_callable(
int32_t callable_id, const void *orch_so_data, size_t orch_so_size, const char *func_name, const char *config_name,
std::vector<std::pair<int, uint64_t>> kernel_addrs, std::vector<ArgDirection> signature
) {
Expand All @@ -645,36 +645,34 @@ int DeviceRunner::register_prepared_callable(
// callable_id; rejecting an out-of-range id here keeps the host and
// AICPU sides in sync and avoids an OOB access at run time.
if (callable_id < 0 || callable_id >= MAX_REGISTERED_CALLABLE_IDS) {
LOG_ERROR(
"register_prepared_callable: callable_id=%d out of range [0, %d)", callable_id, MAX_REGISTERED_CALLABLE_IDS
);
LOG_ERROR("register_callable: callable_id=%d out of range [0, %d)", callable_id, MAX_REGISTERED_CALLABLE_IDS);
return -1;
}
if (orch_so_data == nullptr || orch_so_size == 0) {
LOG_ERROR("register_prepared_callable: empty orch SO for callable_id=%d", callable_id);
LOG_ERROR("register_callable: empty orch SO for callable_id=%d", callable_id);
return -1;
}
if (prepared_callables_.count(callable_id) != 0) {
LOG_ERROR("register_prepared_callable: callable_id=%d already registered", callable_id);
if (callables_.count(callable_id) != 0) {
LOG_ERROR("register_callable: callable_id=%d already registered", callable_id);
return -1;
}

const uint64_t hash = simpler::common::utils::elf_build_id_64(orch_so_data, orch_so_size);

// Hash dedup: share device buffer across callable_ids that carry the same
// SO bytes. Refcount drops in unregister_prepared_callable; we only free
// SO bytes. Refcount drops in unregister_callable; we only free
// when the count hits zero.
auto buf_it = orch_so_dedup_.find(hash);
uint64_t dev_addr = 0;
if (buf_it == orch_so_dedup_.end()) {
void *buf = mem_alloc_.alloc(orch_so_size);
if (buf == nullptr) {
LOG_ERROR("register_prepared_callable: alloc %zu bytes failed", orch_so_size);
LOG_ERROR("register_callable: alloc %zu bytes failed", orch_so_size);
return -1;
}
int rc = rtMemcpy(buf, orch_so_size, orch_so_data, orch_so_size, RT_MEMCPY_HOST_TO_DEVICE);
if (rc != 0) {
LOG_ERROR("register_prepared_callable: rtMemcpy failed: %d", rc);
LOG_ERROR("register_callable: rtMemcpy failed: %d", rc);
mem_alloc_.free(buf);
return rc;
}
Expand All @@ -684,65 +682,62 @@ int DeviceRunner::register_prepared_callable(
entry.refcount = 1;
orch_so_dedup_.emplace(hash, entry);
dev_addr = reinterpret_cast<uint64_t>(buf);
LOG_INFO_V0("register_prepared_callable: hash=0x%lx new buffer %zu bytes", hash, orch_so_size);
LOG_INFO_V0("register_callable: hash=0x%lx new buffer %zu bytes", hash, orch_so_size);
} else {
buf_it->second.refcount++;
dev_addr = reinterpret_cast<uint64_t>(buf_it->second.dev_addr);
LOG_INFO_V0(
"register_prepared_callable: hash=0x%lx shared buffer (refcount=%d)", hash, buf_it->second.refcount
);
LOG_INFO_V0("register_callable: hash=0x%lx shared buffer (refcount=%d)", hash, buf_it->second.refcount);
}

PreparedCallableState state;
CallableState state;
state.hash = hash;
state.dev_orch_so_addr = dev_addr;
state.dev_orch_so_size = orch_so_size;
state.func_name = (func_name != nullptr) ? func_name : "";
state.config_name = (config_name != nullptr) ? config_name : "";
state.kernel_addrs = std::move(kernel_addrs);
state.signature = std::move(signature);
prepared_callables_.emplace(callable_id, std::move(state));
callables_.emplace(callable_id, std::move(state));
return 0;
}

int DeviceRunner::register_prepared_callable_host_orch(
int DeviceRunner::register_callable_host_orch(
int32_t callable_id, void *host_dlopen_handle, void *host_orch_func_ptr,
std::vector<std::pair<int, uint64_t>> kernel_addrs, std::vector<ArgDirection> signature
) {
if (callable_id < 0 || callable_id >= MAX_REGISTERED_CALLABLE_IDS) {
LOG_ERROR(
"register_prepared_callable_host_orch: callable_id=%d out of range [0, %d)", callable_id,
MAX_REGISTERED_CALLABLE_IDS
"register_callable_host_orch: callable_id=%d out of range [0, %d)", callable_id, MAX_REGISTERED_CALLABLE_IDS
);
return -1;
}
if (host_dlopen_handle == nullptr || host_orch_func_ptr == nullptr) {
LOG_ERROR("register_prepared_callable_host_orch: null handle/fn for callable_id=%d", callable_id);
LOG_ERROR("register_callable_host_orch: null handle/fn for callable_id=%d", callable_id);
return -1;
}
if (prepared_callables_.count(callable_id) != 0) {
LOG_ERROR("register_prepared_callable_host_orch: callable_id=%d already registered", callable_id);
if (callables_.count(callable_id) != 0) {
LOG_ERROR("register_callable_host_orch: callable_id=%d already registered", callable_id);
return -1;
}

PreparedCallableState state;
CallableState state;
state.host_dlopen_handle = host_dlopen_handle;
state.host_orch_func_ptr = host_orch_func_ptr;
state.kernel_addrs = std::move(kernel_addrs);
state.signature = std::move(signature);
prepared_callables_.emplace(callable_id, std::move(state));
callables_.emplace(callable_id, std::move(state));
++host_dlopen_total_;
LOG_INFO_V0("register_prepared_callable_host_orch: cid=%d (host dlopen #%zu)", callable_id, host_dlopen_total_);
LOG_INFO_V0("register_callable_host_orch: cid=%d (host dlopen #%zu)", callable_id, host_dlopen_total_);
return 0;
}

int DeviceRunner::unregister_prepared_callable(int32_t callable_id) {
auto it = prepared_callables_.find(callable_id);
if (it == prepared_callables_.end()) {
int DeviceRunner::unregister_callable(int32_t callable_id) {
auto it = callables_.find(callable_id);
if (it == callables_.end()) {
return 0;
}
PreparedCallableState state = std::move(it->second);
prepared_callables_.erase(it);
CallableState state = std::move(it->second);
callables_.erase(it);
aicpu_seen_callable_ids_.erase(callable_id);

if (state.host_dlopen_handle != nullptr) {
Expand All @@ -761,14 +756,12 @@ int DeviceRunner::unregister_prepared_callable(int32_t callable_id) {
return 0;
}

bool DeviceRunner::has_prepared_callable(int32_t callable_id) const {
return prepared_callables_.count(callable_id) != 0;
}
bool DeviceRunner::has_callable(int32_t callable_id) const { return callables_.count(callable_id) != 0; }

BindPreparedCallableResult DeviceRunner::bind_prepared_callable_to_runtime(Runtime &runtime, int32_t callable_id) {
auto it = prepared_callables_.find(callable_id);
if (it == prepared_callables_.end()) {
LOG_ERROR("bind_prepared_callable_to_runtime: callable_id=%d not registered", callable_id);
BindCallableResult DeviceRunner::bind_callable_to_runtime(Runtime &runtime, int32_t callable_id) {
auto it = callables_.find(callable_id);
if (it == callables_.end()) {
LOG_ERROR("bind_callable_to_runtime: callable_id=%d not registered", callable_id);
return {-1, nullptr, nullptr, 0};
}
const auto &state = it->second;
Expand All @@ -779,7 +772,7 @@ BindPreparedCallableResult DeviceRunner::bind_prepared_callable_to_runtime(Runti
// free kernel binaries — but prepared kernels must survive across runs.
for (const auto &kv : state.kernel_addrs) {
if (kv.first < 0 || kv.first >= RUNTIME_MAX_FUNC_ID) {
LOG_ERROR("bind_prepared_callable_to_runtime: func_id=%d out of range", kv.first);
LOG_ERROR("bind_callable_to_runtime: func_id=%d out of range", kv.first);
return {-1, nullptr, nullptr, 0};
}
runtime.replay_function_bin_addr(kv.first, kv.second);
Expand All @@ -790,7 +783,7 @@ BindPreparedCallableResult DeviceRunner::bind_prepared_callable_to_runtime(Runti
// with the authoritative first_sighting answer right before launch.
runtime.set_active_callable_id(callable_id, /*is_new=*/false);
// hbg path: host_orch_func_ptr travels back to the c_api caller, which
// hands it to bind_prepared_to_runtime_impl. trb path: stays null and
// hands it to bind_callable_to_runtime_impl. trb path: stays null and
// the device-side orch SO is resolved from the symbol names above.
return {
0, state.host_orch_func_ptr, state.signature.empty() ? nullptr : state.signature.data(),
Expand Down Expand Up @@ -861,12 +854,12 @@ int DeviceRunner::finalize() {
// each callable_id, so without this loop the host process leaks one
// dlopen handle per (re)created Worker — observable in long-running
// pytest sessions.
for (auto &kv : prepared_callables_) {
for (auto &kv : callables_) {
if (kv.second.host_dlopen_handle != nullptr) {
dlclose(kv.second.host_dlopen_handle);
}
}
prepared_callables_.clear();
callables_.clear();
aicpu_seen_callable_ids_.clear();
aicpu_dlopen_total_ = 0;

Expand Down
38 changes: 19 additions & 19 deletions src/a2a3/platform/onboard/host/device_runner.h
Original file line number Diff line number Diff line change
Expand Up @@ -260,22 +260,22 @@ class DeviceRunner : public DeviceRunnerBase {
* them onto a fresh Runtime without re-uploading.
* @return 0 on success, negative on failure.
*/
int register_prepared_callable(
int register_callable(
int32_t callable_id, const void *orch_so_data, size_t orch_so_size, const char *func_name,
const char *config_name, std::vector<std::pair<int, uint64_t>> kernel_addrs, std::vector<ArgDirection> signature
);

/**
* Host-orchestration variant of register_prepared_callable: stores a
* Host-orchestration variant of register_callable: stores a
* dlopen handle + entry-symbol pointer that runtime_maker resolved on the
* host (host_build_graph variant). Mutually exclusive with the trb-shaped
* `register_prepared_callable` overload — exactly one is invoked for a
* `register_callable` overload — exactly one is invoked for a
* given callable_id, picked by the C ABI based on which staging fields the
* runtime carries after prepare_callable_impl. dlopen handle is owned by
* DeviceRunner from this call onward and dlclose'd by
* unregister_prepared_callable. Increments `host_dlopen_count_`.
* unregister_callable. Increments `host_dlopen_count_`.
*/
int register_prepared_callable_host_orch(
int register_callable_host_orch(
int32_t callable_id, void *host_dlopen_handle, void *host_orch_func_ptr,
std::vector<std::pair<int, uint64_t>> kernel_addrs, std::vector<ArgDirection> signature
);
Expand All @@ -287,17 +287,17 @@ class DeviceRunner : public DeviceRunnerBase {
* callables and only released by finalize().
*
* @param callable_id Id previously passed to one of the
* register_prepared_callable* overloads.
* register_callable* overloads.
* @return 0 on success or if the id was not registered.
*/
int unregister_prepared_callable(int32_t callable_id);
int unregister_callable(int32_t callable_id);

/**
* True iff `callable_id` has prepared state staged via
* register_prepared_callable. Lets the c_api layer reject `run_prepared`
* register_callable. Lets the c_api layer reject `run_prepared`
* calls without a matching `prepare_callable`.
*/
bool has_prepared_callable(int32_t callable_id) const;
bool has_callable(int32_t callable_id) const;

/**
* Replay the prepared state for `callable_id` onto a freshly-constructed
Expand All @@ -306,7 +306,7 @@ class DeviceRunner : public DeviceRunnerBase {
* subsequent `run` dispatches via the AICPU per-cid table. The kernel
* addresses are written directly into func_id_to_addr_ (bypassing
* registered_kernel_func_ids_) so validate_runtime_impl will not free them
* — they survive until unregister_prepared_callable / finalize().
* — they survive until unregister_callable / finalize().
*
* Marks the cid as seen so the upcoming prepare_orch_so resolves
* `register_new_callable_id_` correctly (true exactly on first sighting
Expand All @@ -318,10 +318,10 @@ class DeviceRunner : public DeviceRunnerBase {
* Replay a previously-registered callable's state onto a fresh Runtime
* for a per-run binding. Writes back kernel addrs, orch entry-symbol
* names, and active_callable_id; returns the hbg `host_orch_func_ptr`
* (or nullptr on trb / on error) inside a `BindPreparedCallableResult`
* (or nullptr on trb / on error) inside a `BindCallableResult`
* so the caller can destructure with structured bindings.
*/
BindPreparedCallableResult bind_prepared_callable_to_runtime(Runtime &runtime, int32_t callable_id);
BindCallableResult bind_callable_to_runtime(Runtime &runtime, int32_t callable_id);

/**
* Number of distinct callable_ids the AICPU has been asked to dlopen for.
Expand All @@ -335,7 +335,7 @@ class DeviceRunner : public DeviceRunnerBase {

/**
* Number of host-side dlopen() invocations triggered by
* `register_prepared_callable_host_orch`. Mirrors `aicpu_dlopen_count` but
* `register_callable_host_orch`. Mirrors `aicpu_dlopen_count` but
* counts the host_build_graph variant's host-side dlopens; it never
* decrements (re-prepare after unregister still counts). Tests assert
* `host_dlopen_count == distinct_registered_cids` to verify the prepared
Expand Down Expand Up @@ -372,14 +372,14 @@ class DeviceRunner : public DeviceRunnerBase {

// Per-callable_id prepared state.
//
// `prepared_callables_` maps the caller-stable callable_id to the orch
// `callables_` maps the caller-stable callable_id to the orch
// SO slice + symbol names needed to launch it. `orch_so_dedup_` shares
// device buffers across callable_ids whose orch SO bytes have the same
// ELF Build-ID hash (refcounted; freed when the count hits zero).
// `aicpu_seen_callable_ids_` tracks which ids have already been delivered
// to the AICPU at least once so prepare_orch_so can set
// register_new_callable_id_ correctly on first sighting.
struct PreparedCallableState {
struct CallableState {
// trb path (AICPU dlopens orch SO from device buffer)
uint64_t hash{0};
uint64_t dev_orch_so_addr{0};
Expand All @@ -398,7 +398,7 @@ class DeviceRunner : public DeviceRunnerBase {
size_t capacity{0};
int refcount{0};
};
std::unordered_map<int32_t, PreparedCallableState> prepared_callables_;
std::unordered_map<int32_t, CallableState> callables_;
std::unordered_map<uint64_t, OrchSoBuffer> orch_so_dedup_;
std::unordered_set<int32_t> aicpu_seen_callable_ids_;
// Monotonic count of AICPU dlopens triggered (incremented on each
Expand All @@ -407,7 +407,7 @@ class DeviceRunner : public DeviceRunnerBase {
// re-prepared. Exposed via aicpu_dlopen_count() for tests.
size_t aicpu_dlopen_total_{0};
// Monotonic count of host-side dlopens triggered (incremented on every
// register_prepared_callable_host_orch call; never decremented). Same
// register_callable_host_orch call; never decremented). Same
// re-prepare semantics as aicpu_dlopen_total_, but for hbg variants.
size_t host_dlopen_total_{0};
// ACL lifecycle (process-wide). aclInit must run exactly once; ensure_acl_ready
Expand All @@ -431,8 +431,8 @@ class DeviceRunner : public DeviceRunnerBase {

/**
* Stamp `runtime.{dev_orch_so_addr_, dev_orch_so_size_}` from the
* PreparedCallableState for `runtime.get_active_callable_id()`. The orch
* SO bytes were already H2D'd at `register_prepared_callable` time and
* CallableState for `runtime.get_active_callable_id()`. The orch
* SO bytes were already H2D'd at `register_callable` time and
* are shared via `orch_so_dedup_` across cids; this method only refreshes
* the device-SO metadata onto the per-run Runtime and bumps the AICPU
* first-sighting counter when the cid is new since registration.
Expand Down
Loading
Loading