[xlang] Cross language Python support (#6709)

This commit is contained in:
fyrestone
2020-02-08 13:01:28 +08:00
committed by GitHub
parent f146d05b36
commit 0648bd28ef
59 changed files with 1412 additions and 580 deletions
+35 -25
View File
@@ -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;
}
+1 -1
View File
@@ -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);
+21 -9
View File
@@ -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));
+8 -4
View File
@@ -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)));