mirror of
https://github.com/wassname/ray.git
synced 2026-08-20 12:40:44 +08:00
[Java] Support concurrent actor calls API. (#7022)
* WIP Temp change Attach native thread to jvm * Fix run mode * Address comments.
This commit is contained in:
@@ -961,9 +961,9 @@ Status CoreWorker::ExecuteTask(const TaskSpecification &task_spec,
|
||||
task_type = TaskType::ACTOR_TASK;
|
||||
}
|
||||
|
||||
status = task_execution_callback_(task_type, func,
|
||||
task_spec.GetRequiredResources().GetResourceMap(),
|
||||
args, arg_reference_ids, return_ids, return_objects);
|
||||
status = task_execution_callback_(
|
||||
task_type, func, task_spec.GetRequiredResources().GetResourceMap(), args,
|
||||
arg_reference_ids, return_ids, return_objects, worker_context_.GetWorkerID());
|
||||
|
||||
for (size_t i = 0; i < return_objects->size(); i++) {
|
||||
// The object is nullptr if it already existed in the object store.
|
||||
|
||||
@@ -47,7 +47,7 @@ class CoreWorker : public rpc::CoreWorkerServiceHandler {
|
||||
const std::vector<std::shared_ptr<RayObject>> &args,
|
||||
const std::vector<ObjectID> &arg_reference_ids,
|
||||
const std::vector<ObjectID> &return_ids,
|
||||
std::vector<std::shared_ptr<RayObject>> *results)>;
|
||||
std::vector<std::shared_ptr<RayObject>> *results, const ray::WorkerID &worker_id)>;
|
||||
|
||||
public:
|
||||
/// Construct a CoreWorker instance.
|
||||
|
||||
@@ -56,6 +56,7 @@ jfieldID java_actor_creation_options_default_use_direct_call;
|
||||
jfieldID java_actor_creation_options_max_reconstructions;
|
||||
jfieldID java_actor_creation_options_use_direct_call;
|
||||
jfieldID java_actor_creation_options_jvm_options;
|
||||
jfieldID java_actor_creation_options_max_concurrency;
|
||||
|
||||
jclass java_gcs_client_options_class;
|
||||
jfieldID java_gcs_client_options_ip;
|
||||
@@ -69,6 +70,7 @@ jfieldID java_native_ray_object_metadata;
|
||||
|
||||
jclass java_task_executor_class;
|
||||
jmethodID java_task_executor_execute;
|
||||
jmethodID java_task_executor_get;
|
||||
|
||||
JavaVM *jvm;
|
||||
|
||||
@@ -164,7 +166,8 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
|
||||
env->GetFieldID(java_actor_creation_options_class, "useDirectCall", "Z");
|
||||
java_actor_creation_options_jvm_options = env->GetFieldID(
|
||||
java_actor_creation_options_class, "jvmOptions", "Ljava/lang/String;");
|
||||
|
||||
java_actor_creation_options_max_concurrency =
|
||||
env->GetFieldID(java_actor_creation_options_class, "maxConcurrency", "I");
|
||||
java_gcs_client_options_class = LoadClass(env, "org/ray/runtime/gcs/GcsClientOptions");
|
||||
java_gcs_client_options_ip =
|
||||
env->GetFieldID(java_gcs_client_options_class, "ip", "Ljava/lang/String;");
|
||||
@@ -186,6 +189,11 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
|
||||
env->GetMethodID(java_task_executor_class, "execute",
|
||||
"(Ljava/util/List;Ljava/util/List;)Ljava/util/List;");
|
||||
|
||||
java_task_executor_get = env->GetStaticMethodID(
|
||||
java_task_executor_class,
|
||||
"get",
|
||||
"([B)Lorg/ray/runtime/task/TaskExecutor;");
|
||||
|
||||
return CURRENT_JNI_VERSION;
|
||||
}
|
||||
|
||||
|
||||
@@ -105,6 +105,8 @@ extern jfieldID java_actor_creation_options_max_reconstructions;
|
||||
extern jfieldID java_actor_creation_options_use_direct_call;
|
||||
/// jvmOptions field of ActorCreationOptions class
|
||||
extern jfieldID java_actor_creation_options_jvm_options;
|
||||
/// maxConcurrency field of ActorCreationOptions class
|
||||
extern jfieldID java_actor_creation_options_max_concurrency;
|
||||
|
||||
/// GcsClientOptions class
|
||||
extern jclass java_gcs_client_options_class;
|
||||
@@ -129,6 +131,9 @@ extern jclass java_task_executor_class;
|
||||
/// execute method of TaskExecutor class
|
||||
extern jmethodID java_task_executor_execute;
|
||||
|
||||
/// The `get` method in TaskExecutor class
|
||||
extern jmethodID java_task_executor_get;
|
||||
|
||||
#define CURRENT_JNI_VERSION JNI_VERSION_1_8
|
||||
|
||||
extern JavaVM *jvm;
|
||||
|
||||
@@ -6,7 +6,6 @@
|
||||
#include "ray/core_worker/lib/java/jni_utils.h"
|
||||
|
||||
thread_local JNIEnv *local_env = nullptr;
|
||||
thread_local jobject local_java_task_executor = nullptr;
|
||||
|
||||
inline ray::gcs::GcsClientOptions ToGcsClientOptions(JNIEnv *env,
|
||||
jobject gcs_client_options) {
|
||||
@@ -39,9 +38,23 @@ JNIEXPORT jlong JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeInitCoreWork
|
||||
const std::vector<std::shared_ptr<ray::RayObject>> &args,
|
||||
const std::vector<ObjectID> &arg_reference_ids,
|
||||
const std::vector<ObjectID> &return_ids,
|
||||
std::vector<std::shared_ptr<ray::RayObject>> *results) {
|
||||
std::vector<std::shared_ptr<ray::RayObject>> *results,
|
||||
const ray::WorkerID &worker_id) {
|
||||
JNIEnv *env = local_env;
|
||||
if (!env) {
|
||||
// Attach the native thread to JVM.
|
||||
auto status =
|
||||
jvm->AttachCurrentThreadAsDaemon(reinterpret_cast<void **>(&env), nullptr);
|
||||
RAY_CHECK(status == JNI_OK) << "Failed to get JNIEnv. Return code: " << status;
|
||||
local_env = env;
|
||||
}
|
||||
|
||||
RAY_CHECK(env);
|
||||
|
||||
auto worker_id_bytes = IdToJavaByteArray<ray::WorkerID>(env, worker_id);
|
||||
jobject local_java_task_executor = env->CallStaticObjectMethod(
|
||||
java_task_executor_class, java_task_executor_get, worker_id_bytes);
|
||||
|
||||
RAY_CHECK(local_java_task_executor);
|
||||
// convert RayFunction
|
||||
jobject ray_function_array_list = NativeRayFunctionDescriptorToJavaStringList(
|
||||
@@ -87,13 +100,11 @@ JNIEXPORT jlong JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeInitCoreWork
|
||||
}
|
||||
|
||||
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeRunTaskExecutor(
|
||||
JNIEnv *env, jclass o, jlong nativeCoreWorkerPointer, jobject javaTaskExecutor) {
|
||||
JNIEnv *env, jclass o, jlong nativeCoreWorkerPointer) {
|
||||
local_env = env;
|
||||
local_java_task_executor = javaTaskExecutor;
|
||||
auto core_worker = reinterpret_cast<ray::CoreWorker *>(nativeCoreWorkerPointer);
|
||||
core_worker->StartExecutingTasks();
|
||||
local_env = nullptr;
|
||||
local_java_task_executor = nullptr;
|
||||
}
|
||||
|
||||
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeDestroyCoreWorker(
|
||||
|
||||
@@ -19,10 +19,10 @@ JNIEXPORT jlong JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeInitCoreWork
|
||||
/*
|
||||
* Class: org_ray_runtime_RayNativeRuntime
|
||||
* Method: nativeRunTaskExecutor
|
||||
* Signature: (JLorg/ray/runtime/task/TaskExecutor;)V
|
||||
* Signature: (J)V
|
||||
*/
|
||||
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeRunTaskExecutor(
|
||||
JNIEnv *, jclass, jlong, jobject);
|
||||
JNIEnv *, jclass, jlong);
|
||||
|
||||
/*
|
||||
* Class: org_ray_runtime_RayNativeRuntime
|
||||
|
||||
@@ -92,6 +92,7 @@ inline ray::ActorCreationOptions ToActorCreationOptions(JNIEnv *env,
|
||||
bool use_direct_call;
|
||||
std::unordered_map<std::string, double> resources;
|
||||
std::vector<std::string> dynamic_worker_options;
|
||||
uint64_t max_concurrency = 1;
|
||||
if (actorCreationOptions) {
|
||||
max_reconstructions = static_cast<uint64_t>(env->GetIntField(
|
||||
actorCreationOptions, java_actor_creation_options_max_reconstructions));
|
||||
@@ -106,6 +107,8 @@ inline ray::ActorCreationOptions ToActorCreationOptions(JNIEnv *env,
|
||||
std::string jvm_options = JavaStringToNativeString(env, java_jvm_options);
|
||||
dynamic_worker_options.emplace_back(jvm_options);
|
||||
}
|
||||
max_concurrency = static_cast<uint64_t>(env->GetIntField(
|
||||
actorCreationOptions, java_actor_creation_options_max_concurrency));
|
||||
} else {
|
||||
use_direct_call =
|
||||
env->GetStaticBooleanField(java_actor_creation_options_class,
|
||||
@@ -115,7 +118,7 @@ inline ray::ActorCreationOptions ToActorCreationOptions(JNIEnv *env,
|
||||
ray::ActorCreationOptions actor_creation_options{
|
||||
static_cast<uint64_t>(max_reconstructions),
|
||||
use_direct_call,
|
||||
/*max_concurrency=*/1,
|
||||
static_cast<int>(max_concurrency),
|
||||
resources,
|
||||
resources,
|
||||
dynamic_worker_options,
|
||||
|
||||
Reference in New Issue
Block a user