RPC_CMD_HELLO,
RPC_CMD_DEVICE_COUNT,
RPC_CMD_GRAPH_RECOMPUTE,
+ RPC_CMD_MEMSET_TENSOR,
RPC_CMD_COUNT,
};
uint8_t value;
};
+struct rpc_msg_memset_tensor_req {
+ rpc_tensor tensor;
+ uint64_t offset;
+ uint64_t size;
+ uint8_t value;
+};
+
struct rpc_msg_set_tensor_hash_req {
rpc_tensor tensor;
uint64_t offset;
return GGML_STATUS_SUCCESS;
}
+static void ggml_backend_rpc_buffer_memset_tensor(
+ ggml_backend_buffer_t buffer, ggml_tensor * tensor, uint8_t value, size_t offset, size_t size) {
+ ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context;
+ rpc_msg_memset_tensor_req request = {
+ /* .tensor = */ serialize_tensor(tensor),
+ /* .offset = */ offset,
+ /* .size = */ size,
+ /* .value = */ value,
+ };
+ bool status = send_rpc_cmd(ctx->sock, RPC_CMD_MEMSET_TENSOR, &request, sizeof(request), nullptr, 0);
+ RPC_STATUS_ASSERT(status);
+}
+
static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, size_t offset, size_t size) {
ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context;
rpc_tensor rpc_tensor = serialize_tensor(tensor);
/* .free_buffer = */ ggml_backend_rpc_buffer_free_buffer,
/* .get_base = */ ggml_backend_rpc_buffer_get_base,
/* .init_tensor = */ ggml_backend_rpc_buffer_init_tensor,
- /* .memset_tensor = */ NULL,
+ /* .memset_tensor = */ ggml_backend_rpc_buffer_memset_tensor,
/* .set_tensor = */ ggml_backend_rpc_buffer_set_tensor,
/* .get_tensor = */ ggml_backend_rpc_buffer_get_tensor,
/* .set_tensor_2d = */ NULL,
bool buffer_get_base(const rpc_msg_buffer_get_base_req & request, rpc_msg_buffer_get_base_rsp & response);
bool free_buffer(const rpc_msg_free_buffer_req & request);
bool buffer_clear(const rpc_msg_buffer_clear_req & request);
+ bool memset_tensor(const rpc_msg_memset_tensor_req & request);
bool set_tensor(const std::vector<uint8_t> & input);
bool set_tensor_hash(const rpc_msg_set_tensor_hash_req & request, rpc_msg_set_tensor_hash_rsp & response);
bool get_tensor(const rpc_msg_get_tensor_req & request, std::vector<uint8_t> & response);
return true;
}
+bool rpc_server::memset_tensor(const rpc_msg_memset_tensor_req & request) {
+ struct ggml_init_params params {
+ /*.mem_size =*/ ggml_tensor_overhead(),
+ /*.mem_buffer =*/ NULL,
+ /*.no_alloc =*/ true,
+ };
+ ggml_context_ptr ctx_ptr { ggml_init(params) };
+ GGML_ASSERT(ctx_ptr != nullptr);
+ ggml_context * ctx = ctx_ptr.get();
+ ggml_tensor * tensor = deserialize_tensor(ctx, &request.tensor);
+ if (tensor == nullptr || tensor->buffer == nullptr) {
+ GGML_LOG_ERROR("[%s] error deserializing tensor\n", __func__);
+ return false;
+ }
+
+ const uint64_t tensor_size = ggml_nbytes(tensor);
+ if (request.offset > tensor_size || request.size > tensor_size - request.offset) {
+ GGML_LOG_ERROR("[%s] tensor region (offset=%" PRIu64 ", size=%" PRIu64 ") out of tensor bounds [0, %" PRIu64 ")\n",
+ __func__, request.offset, request.size, tensor_size);
+ return false;
+ }
+
+ const uint64_t buffer_start = (uint64_t) ggml_backend_buffer_get_base(tensor->buffer);
+ const uint64_t buffer_size = ggml_backend_buffer_get_size(tensor->buffer);
+ if (request.tensor.data < buffer_start) {
+ GGML_LOG_ERROR("[%s] tensor data before buffer start\n", __func__);
+ return false;
+ }
+ const uint64_t data_offset = request.tensor.data - buffer_start;
+ if (data_offset > buffer_size ||
+ request.offset > buffer_size - data_offset ||
+ request.size > buffer_size - data_offset - request.offset) {
+ GGML_LOG_ERROR("[%s] tensor region out of buffer bounds\n", __func__);
+ return false;
+ }
+ if (tensor->buffer->iface.memset_tensor == nullptr) {
+ GGML_LOG_ERROR("[%s] memset not implemented by backend buffer\n", __func__);
+ return false;
+ }
+
+ LOG_DBG("[%s] buffer: %p, data: %p, offset: %" PRIu64 ", size: %" PRIu64 ", value: %u\n",
+ __func__, (void *) tensor->buffer, tensor->data, request.offset, request.size, request.value);
+ ggml_backend_tensor_memset(tensor, request.value, request.offset, request.size);
+ return true;
+}
+
ggml_tensor * rpc_server::deserialize_tensor(struct ggml_context * ctx, const rpc_tensor * tensor) {
// Validate tensor type before using it
if (tensor->type >= GGML_TYPE_COUNT) {
}
break;
}
+ case RPC_CMD_MEMSET_TENSOR: {
+ rpc_msg_memset_tensor_req request;
+ if (!recv_msg(sock, &request, sizeof(request))) {
+ return;
+ }
+ if (!server.memset_tensor(request)) {
+ return;
+ }
+ if (!send_msg(sock, nullptr, 0)) {
+ return;
+ }
+ break;
+ }
case RPC_CMD_SET_TENSOR: {
std::vector<uint8_t> input;
if (!recv_msg(sock, input)) {