mirror of
https://github.com/wassname/ray.git
synced 2026-08-14 12:40:23 +08:00
[xlang] Cross language Python support (#6709)
This commit is contained in:
@@ -1,27 +1,25 @@
|
||||
#include "streaming_jni_common.h"
|
||||
|
||||
std::vector<ray::ObjectID>
|
||||
jarray_to_object_id_vec(JNIEnv *env, jobjectArray jarr) {
|
||||
std::vector<ray::ObjectID> jarray_to_object_id_vec(JNIEnv *env, jobjectArray jarr) {
|
||||
int stringCount = env->GetArrayLength(jarr);
|
||||
std::vector<ray::ObjectID> object_id_vec;
|
||||
for (int i = 0; i < stringCount; i++) {
|
||||
auto jstr = (jbyteArray) (env->GetObjectArrayElement(jarr, i));
|
||||
auto jstr = (jbyteArray)(env->GetObjectArrayElement(jarr, i));
|
||||
UniqueIdFromJByteArray idFromJByteArray(env, jstr);
|
||||
object_id_vec.push_back(idFromJByteArray.PID);
|
||||
}
|
||||
return object_id_vec;
|
||||
return object_id_vec;
|
||||
}
|
||||
|
||||
std::vector<ray::ActorID>
|
||||
jarray_to_actor_id_vec(JNIEnv *env, jobjectArray jarr) {
|
||||
std::vector<ray::ActorID> jarray_to_actor_id_vec(JNIEnv *env, jobjectArray jarr) {
|
||||
int count = env->GetArrayLength(jarr);
|
||||
std::vector<ray::ActorID> actor_id_vec;
|
||||
for (int i = 0; i < count; i++) {
|
||||
auto bytes = (jbyteArray)(env->GetObjectArrayElement(jarr, i));
|
||||
std::string id_str(ray::ActorID::Size(), 0);
|
||||
env->GetByteArrayRegion(bytes, 0, ray::ActorID::Size(),
|
||||
reinterpret_cast<jbyte *>(&id_str.front()));
|
||||
actor_id_vec.push_back(ActorID::FromBinary(id_str));
|
||||
reinterpret_cast<jbyte *>(&id_str.front()));
|
||||
actor_id_vec.push_back(ActorID::FromBinary(id_str));
|
||||
}
|
||||
|
||||
return actor_id_vec;
|
||||
@@ -38,17 +36,22 @@ jint throwChannelInitException(JNIEnv *env, const char *message,
|
||||
const std::vector<ray::ObjectID> &abnormal_queues) {
|
||||
jclass array_list_class = env->FindClass("java/util/ArrayList");
|
||||
jmethodID array_list_constructor = env->GetMethodID(array_list_class, "<init>", "()V");
|
||||
jmethodID array_list_add = env->GetMethodID(array_list_class, "add", "(Ljava/lang/Object;)Z");
|
||||
jmethodID array_list_add =
|
||||
env->GetMethodID(array_list_class, "add", "(Ljava/lang/Object;)Z");
|
||||
jobject array_list = env->NewObject(array_list_class, array_list_constructor);
|
||||
|
||||
for (auto &q_id : abnormal_queues) {
|
||||
jbyteArray jbyte_array = env->NewByteArray(kUniqueIDSize);
|
||||
env->SetByteArrayRegion(jbyte_array, 0, kUniqueIDSize, const_cast<jbyte*>(reinterpret_cast<const jbyte *>(q_id.Data())));
|
||||
env->SetByteArrayRegion(
|
||||
jbyte_array, 0, kUniqueIDSize,
|
||||
const_cast<jbyte *>(reinterpret_cast<const jbyte *>(q_id.Data())));
|
||||
env->CallBooleanMethod(array_list, array_list_add, jbyte_array);
|
||||
}
|
||||
|
||||
jclass ex_class = env->FindClass("org/ray/streaming/runtime/transfer/ChannelInitException");
|
||||
jmethodID ex_constructor = env->GetMethodID(ex_class, "<init>", "(Ljava/lang/String;Ljava/util/List;)V");
|
||||
jclass ex_class =
|
||||
env->FindClass("org/ray/streaming/runtime/transfer/ChannelInitException");
|
||||
jmethodID ex_constructor =
|
||||
env->GetMethodID(ex_class, "<init>", "(Ljava/lang/String;Ljava/util/List;)V");
|
||||
jstring message_jstr = env->NewStringUTF(message);
|
||||
jobject ex_obj = env->NewObject(ex_class, ex_constructor, message_jstr, array_list);
|
||||
env->DeleteLocalRef(message_jstr);
|
||||
@@ -56,7 +59,8 @@ jint throwChannelInitException(JNIEnv *env, const char *message,
|
||||
}
|
||||
|
||||
jint throwChannelInterruptException(JNIEnv *env, const char *message) {
|
||||
jclass ex_class = env->FindClass("org/ray/streaming/runtime/transfer/ChannelInterruptException");
|
||||
jclass ex_class =
|
||||
env->FindClass("org/ray/streaming/runtime/transfer/ChannelInterruptException");
|
||||
return env->ThrowNew(ex_class, message);
|
||||
}
|
||||
|
||||
@@ -69,12 +73,13 @@ jclass LoadClass(JNIEnv *env, const char *class_name) {
|
||||
}
|
||||
|
||||
template <typename NativeT>
|
||||
void JavaListToNativeVector(
|
||||
JNIEnv *env, jobject java_list, std::vector<NativeT> *native_vector,
|
||||
std::function<NativeT(JNIEnv *, jobject)> element_converter) {
|
||||
void JavaListToNativeVector(JNIEnv *env, jobject java_list,
|
||||
std::vector<NativeT> *native_vector,
|
||||
std::function<NativeT(JNIEnv *, jobject)> element_converter) {
|
||||
jclass java_list_class = LoadClass(env, "java/util/List");
|
||||
jmethodID java_list_size = env->GetMethodID(java_list_class, "size", "()I");
|
||||
jmethodID java_list_get = env->GetMethodID(java_list_class, "get", "(I)Ljava/lang/Object;");
|
||||
jmethodID java_list_get =
|
||||
env->GetMethodID(java_list_class, "get", "(I)Ljava/lang/Object;");
|
||||
int size = env->CallIntMethod(java_list, java_list_size);
|
||||
native_vector->clear();
|
||||
for (int i = 0; i < size; i++) {
|
||||
@@ -100,24 +105,29 @@ void JavaStringListToNativeStringVector(JNIEnv *env, jobject java_list,
|
||||
});
|
||||
}
|
||||
|
||||
ray::RayFunction FunctionDescriptorToRayFunction(JNIEnv *env, jobject functionDescriptor) {
|
||||
jclass java_language_class = LoadClass(env, "org/ray/runtime/generated/Common$Language");
|
||||
ray::RayFunction FunctionDescriptorToRayFunction(JNIEnv *env,
|
||||
jobject functionDescriptor) {
|
||||
jclass java_language_class =
|
||||
LoadClass(env, "org/ray/runtime/generated/Common$Language");
|
||||
jclass java_function_descriptor_class =
|
||||
LoadClass(env, "org/ray/runtime/functionmanager/FunctionDescriptor");
|
||||
jmethodID java_language_get_number = env->GetMethodID(java_language_class, "getNumber", "()I");
|
||||
jmethodID java_language_get_number =
|
||||
env->GetMethodID(java_language_class, "getNumber", "()I");
|
||||
jmethodID java_function_descriptor_get_language =
|
||||
env->GetMethodID(java_function_descriptor_class, "getLanguage",
|
||||
"()Lorg/ray/runtime/generated/Common$Language;");
|
||||
jobject java_language =
|
||||
env->CallObjectMethod(functionDescriptor, java_function_descriptor_get_language);
|
||||
int language = env->CallIntMethod(java_language, java_language_get_number);
|
||||
std::vector<std::string> function_descriptor;
|
||||
auto language = static_cast<::Language>(
|
||||
env->CallIntMethod(java_language, java_language_get_number));
|
||||
std::vector<std::string> function_descriptor_list;
|
||||
jmethodID java_function_descriptor_to_list =
|
||||
env->GetMethodID(java_function_descriptor_class, "toList", "()Ljava/util/List;");
|
||||
JavaStringListToNativeStringVector(
|
||||
env, env->CallObjectMethod(functionDescriptor, java_function_descriptor_to_list),
|
||||
&function_descriptor);
|
||||
ray::RayFunction ray_function{static_cast<::Language>(language), function_descriptor};
|
||||
&function_descriptor_list);
|
||||
ray::FunctionDescriptor function_descriptor =
|
||||
ray::FunctionDescriptorBuilder::FromVector(language, function_descriptor_list);
|
||||
ray::RayFunction ray_function{language, function_descriptor};
|
||||
return ray_function;
|
||||
}
|
||||
|
||||
|
||||
@@ -77,7 +77,7 @@ std::shared_ptr<LocalMemoryBuffer> Transport::SendForResultWithRetry(
|
||||
int64_t timeout_ms) {
|
||||
STREAMING_LOG(INFO) << "SendForResultWithRetry retry_cnt: " << retry_cnt
|
||||
<< " timeout_ms: " << timeout_ms
|
||||
<< " function: " << function.GetFunctionDescriptor()[0];
|
||||
<< " function: " << function.GetFunctionDescriptor()->ToString();
|
||||
std::shared_ptr<LocalMemoryBuffer> buffer_shared = std::move(buffer);
|
||||
for (int cnt = 0; cnt < retry_cnt; cnt++) {
|
||||
auto result = SendForResult(buffer_shared, function, timeout_ms);
|
||||
|
||||
@@ -287,10 +287,18 @@ class StreamingWorker {
|
||||
JobID::FromInt(1), gcs_options, "", "127.0.0.1", node_manager_port,
|
||||
std::bind(&StreamingWorker::ExecuteTask, this, _1, _2, _3, _4, _5, _6, _7));
|
||||
|
||||
RayFunction reader_async_call_func{ray::Language::PYTHON, {"reader_async_call_func"}};
|
||||
RayFunction reader_sync_call_func{ray::Language::PYTHON, {"reader_sync_call_func"}};
|
||||
RayFunction writer_async_call_func{ray::Language::PYTHON, {"writer_async_call_func"}};
|
||||
RayFunction writer_sync_call_func{ray::Language::PYTHON, {"writer_sync_call_func"}};
|
||||
RayFunction reader_async_call_func{ray::Language::PYTHON,
|
||||
ray::FunctionDescriptorBuilder::BuildPython(
|
||||
"reader_async_call_func", "", "", "")};
|
||||
RayFunction reader_sync_call_func{
|
||||
ray::Language::PYTHON,
|
||||
ray::FunctionDescriptorBuilder::BuildPython("reader_sync_call_func", "", "", "")};
|
||||
RayFunction writer_async_call_func{ray::Language::PYTHON,
|
||||
ray::FunctionDescriptorBuilder::BuildPython(
|
||||
"writer_async_call_func", "", "", "")};
|
||||
RayFunction writer_sync_call_func{
|
||||
ray::Language::PYTHON,
|
||||
ray::FunctionDescriptorBuilder::BuildPython("writer_sync_call_func", "", "", "")};
|
||||
|
||||
reader_client_ = std::make_shared<ReaderClient>(worker_.get(), reader_async_call_func,
|
||||
reader_sync_call_func);
|
||||
@@ -314,18 +322,22 @@ class StreamingWorker {
|
||||
// Only one arg param used in streaming.
|
||||
STREAMING_CHECK(args.size() >= 1) << "args.size() = " << args.size();
|
||||
|
||||
std::vector<std::string> function_descriptor = ray_function.GetFunctionDescriptor();
|
||||
STREAMING_LOG(INFO) << "StreamingWorker::ExecuteTask " << function_descriptor[0];
|
||||
ray::FunctionDescriptor function_descriptor = ray_function.GetFunctionDescriptor();
|
||||
RAY_CHECK(function_descriptor->Type() ==
|
||||
ray::FunctionDescriptorType::kPythonFunctionDescriptor);
|
||||
auto typed_descriptor = function_descriptor->As<ray::PythonFunctionDescriptor>();
|
||||
STREAMING_LOG(INFO) << "StreamingWorker::ExecuteTask "
|
||||
<< typed_descriptor->ModuleName();
|
||||
|
||||
std::string func_name = function_descriptor[0];
|
||||
std::string func_name = typed_descriptor->ModuleName();
|
||||
if (func_name == "init") {
|
||||
std::shared_ptr<LocalMemoryBuffer> local_buffer =
|
||||
std::make_shared<LocalMemoryBuffer>(args[0]->GetData()->Data(),
|
||||
args[0]->GetData()->Size(), true);
|
||||
HandleInitTask(local_buffer);
|
||||
} else if (func_name == "execute_test") {
|
||||
STREAMING_LOG(INFO) << "Test name: " << function_descriptor[1];
|
||||
test_suite_->ExecuteTest(function_descriptor[1]);
|
||||
STREAMING_LOG(INFO) << "Test name: " << typed_descriptor->ClassName();
|
||||
test_suite_->ExecuteTest(typed_descriptor->ClassName());
|
||||
} else if (func_name == "check_current_test_status") {
|
||||
results->push_back(
|
||||
std::make_shared<RayObject>(test_suite_->CheckCurTestStatus(), nullptr));
|
||||
|
||||
@@ -162,7 +162,8 @@ class StreamingQueueTestBase : public ::testing::TestWithParam<uint64_t> {
|
||||
std::unordered_map<std::string, double> resources;
|
||||
TaskOptions options{0, true, resources};
|
||||
std::vector<ObjectID> return_ids;
|
||||
RayFunction func{ray::Language::PYTHON, {"init"}};
|
||||
RayFunction func{ray::Language::PYTHON,
|
||||
ray::FunctionDescriptorBuilder::BuildPython("init", "", "", "")};
|
||||
|
||||
RAY_CHECK_OK(driver.SubmitActorTask(self_actor_id, func, args, options, &return_ids));
|
||||
}
|
||||
@@ -176,7 +177,8 @@ class StreamingQueueTestBase : public ::testing::TestWithParam<uint64_t> {
|
||||
std::unordered_map<std::string, double> resources;
|
||||
TaskOptions options{0, true, resources};
|
||||
std::vector<ObjectID> return_ids;
|
||||
RayFunction func{ray::Language::PYTHON, {"execute_test", test}};
|
||||
RayFunction func{ray::Language::PYTHON, ray::FunctionDescriptorBuilder::BuildPython(
|
||||
"execute_test", test, "", "")};
|
||||
|
||||
RAY_CHECK_OK(driver.SubmitActorTask(actor_id, func, args, options, &return_ids));
|
||||
}
|
||||
@@ -190,7 +192,8 @@ class StreamingQueueTestBase : public ::testing::TestWithParam<uint64_t> {
|
||||
std::unordered_map<std::string, double> resources;
|
||||
TaskOptions options{1, true, resources};
|
||||
std::vector<ObjectID> return_ids;
|
||||
RayFunction func{ray::Language::PYTHON, {"check_current_test_status"}};
|
||||
RayFunction func{ray::Language::PYTHON, ray::FunctionDescriptorBuilder::BuildPython(
|
||||
"check_current_test_status", "", "", "")};
|
||||
|
||||
RAY_CHECK_OK(driver.SubmitActorTask(actor_id, func, args, options, &return_ids));
|
||||
|
||||
@@ -250,7 +253,8 @@ class StreamingQueueTestBase : public ::testing::TestWithParam<uint64_t> {
|
||||
uint8_t array[] = {1, 2, 3};
|
||||
auto buffer = std::make_shared<LocalMemoryBuffer>(array, sizeof(array));
|
||||
|
||||
RayFunction func{ray::Language::PYTHON, {"actor creation task"}};
|
||||
RayFunction func{ray::Language::PYTHON, ray::FunctionDescriptorBuilder::BuildPython(
|
||||
"actor creation task", "", "", "")};
|
||||
std::vector<TaskArg> args;
|
||||
args.emplace_back(TaskArg::PassByValue(std::make_shared<RayObject>(buffer, nullptr)));
|
||||
|
||||
|
||||
Reference in New Issue
Block a user