[Java] Local and distributed ref counting in Java (#9371)

This commit is contained in:
Kai Yang
2020-07-31 11:49:31 +08:00
committed by GitHub
parent e2c0174ab2
commit 02fd950252
42 changed files with 1072 additions and 206 deletions
+12
View File
@@ -164,6 +164,18 @@ void CoreWorkerProcess::EnsureInitialized() {
<< "shutdown.";
}
std::shared_ptr<CoreWorker> CoreWorkerProcess::TryGetWorker(const WorkerID &worker_id) {
if (!instance_) {
return nullptr;
}
absl::ReaderMutexLock workers_lock(&instance_->worker_map_mutex_);
auto it = instance_->workers_.find(worker_id);
if (it != instance_->workers_.end()) {
return it->second;
}
return nullptr;
}
CoreWorker &CoreWorkerProcess::GetCoreWorker() {
EnsureInitialized();
if (instance_->options_.num_workers == 1) {
+7
View File
@@ -178,6 +178,13 @@ class CoreWorkerProcess {
/// `CoreWorkerProcess` has full control of the destruction timing of `CoreWorker`.
static CoreWorker &GetCoreWorker();
/// Try to get the `CoreWorker` instance by worker ID.
/// If the current thread is not associated with a core worker, returns a null pointer.
///
/// \param[in] workerId The worker ID.
/// \return The `CoreWorker` instance.
static std::shared_ptr<CoreWorker> TryGetWorker(const WorkerID &worker_id);
/// Set the core worker associated with the current thread by worker ID.
/// Currently used by Java worker only.
///
@@ -18,10 +18,10 @@
#include <sstream>
#include "jni_utils.h"
#include "ray/common/id.h"
#include "ray/core_worker/actor_handle.h"
#include "ray/core_worker/core_worker.h"
#include "jni_utils.h"
thread_local JNIEnv *local_env = nullptr;
jobject java_task_executor = nullptr;
@@ -72,6 +72,19 @@ jobject ToJavaArgs(JNIEnv *env, jbooleanArray java_check_results,
}
}
JNIEnv *GetJNIEnv() {
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);
return env;
}
#ifdef __cplusplus
extern "C" {
#endif
@@ -98,16 +111,7 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
const std::vector<ObjectID> &arg_reference_ids,
const std::vector<ObjectID> &return_ids,
std::vector<std::shared_ptr<ray::RayObject>> *results) {
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);
JNIEnv *env = GetJNIEnv();
RAY_CHECK(java_task_executor);
// convert RayFunction
@@ -141,6 +145,8 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
env->CallObjectMethod(java_task_executor, java_task_executor_execute,
ray_function_array_list, args_array_list);
RAY_CHECK_JAVA_EXCEPTION(env);
// Process return objects.
if (!return_ids.empty()) {
std::vector<std::shared_ptr<ray::RayObject>> return_objects;
JavaListToNativeVector<std::shared_ptr<ray::RayObject>>(
@@ -148,8 +154,26 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
[](JNIEnv *env, jobject java_native_ray_object) {
return JavaNativeRayObjectToNativeRayObject(env, java_native_ray_object);
});
for (auto &obj : return_objects) {
results->push_back(obj);
std::vector<size_t> data_sizes;
std::vector<std::shared_ptr<ray::Buffer>> metadatas;
std::vector<std::vector<ray::ObjectID>> contained_object_ids;
for (size_t i = 0; i < return_objects.size(); i++) {
data_sizes.push_back(
return_objects[i]->HasData() ? return_objects[i]->GetData()->Size() : 0);
metadatas.push_back(return_objects[i]->GetMetadata());
contained_object_ids.push_back(return_objects[i]->GetNestedIds());
}
RAY_CHECK_OK(ray::CoreWorkerProcess::GetCoreWorker().AllocateReturnObjects(
return_ids, data_sizes, metadatas, contained_object_ids, results));
for (size_t i = 0; i < data_sizes.size(); i++) {
auto result = (*results)[i];
// A nullptr is returned if the object already exists.
if (result != nullptr) {
if (result->HasData()) {
memcpy(result->GetData()->Data(), return_objects[i]->GetData()->Data(),
data_sizes[i]);
}
}
}
}
@@ -159,6 +183,26 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
return ray::Status::OK();
};
auto gc_collect = []() {
// A Java worker process usually contains more than one worker.
// A LocalGC request is likely to be received by multiple workers in a short time.
// Here we ensure that the 1 second interval of `System.gc()` execution is
// guaranteed no matter how frequent the requests are received and how many workers
// the process has.
static absl::Mutex mutex;
static int64_t last_gc_time_ms = 0;
absl::MutexLock lock(&mutex);
int64_t start = current_time_ms();
if (last_gc_time_ms + 1000 < start) {
JNIEnv *env = GetJNIEnv();
RAY_LOG(INFO) << "Calling System.gc() ...";
env->CallStaticObjectMethod(java_system_class, java_system_gc);
last_gc_time_ms = current_time_ms();
RAY_LOG(INFO) << "GC finished in " << (double) (last_gc_time_ms - start) / 1000
<< " seconds.";
}
};
ray::CoreWorkerOptions options = {
static_cast<ray::WorkerType>(workerMode), // worker_type
ray::Language::JAVA, // langauge
@@ -178,10 +222,10 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
"", // stderr_file
task_execution_callback, // task_execution_callback
nullptr, // check_signals
nullptr, // gc_collect
gc_collect, // gc_collect
nullptr, // get_lang_stack
nullptr, // kill_main
false, // ref_counting_enabled
true, // ref_counting_enabled
false, // is_local_mode
static_cast<int>(numWorkersPerProcess), // num_workers
};
@@ -14,10 +14,47 @@
#include "io_ray_runtime_object_NativeObjectStore.h"
#include <jni.h>
#include "jni_utils.h"
#include "ray/common/id.h"
#include "ray/core_worker/common.h"
#include "ray/core_worker/core_worker.h"
#include "jni_utils.h"
ray::Status PutSerializedObject(JNIEnv *env, jobject obj, ray::ObjectID object_id,
ray::ObjectID *out_object_id, bool pin_object = true) {
auto native_ray_object = JavaNativeRayObjectToNativeRayObject(env, obj);
RAY_CHECK(native_ray_object != nullptr);
size_t data_size = 0;
if (native_ray_object->HasData()) {
data_size = native_ray_object->GetData()->Size();
}
std::shared_ptr<ray::Buffer> data;
ray::Status status;
if (object_id.IsNil()) {
status = ray::CoreWorkerProcess::GetCoreWorker().Create(
native_ray_object->GetMetadata(), data_size, native_ray_object->GetNestedIds(),
out_object_id, &data);
} else {
status = ray::CoreWorkerProcess::GetCoreWorker().Create(
native_ray_object->GetMetadata(), data_size, object_id, &data);
*out_object_id = object_id;
}
if (!status.ok()) {
return status;
}
// If data is nullptr, that means the ObjectID already existed, which we ignore.
// TODO(edoakes): this is hacky, we should return the error instead and deal with it
// here.
if (data != nullptr) {
if (data->Size() > 0) {
memcpy(data->Data(), native_ray_object->GetData()->Data(), data->Size());
}
RAY_CHECK_OK(ray::CoreWorkerProcess::GetCoreWorker().Seal(
*out_object_id, pin_object && object_id.IsNil()));
}
return ray::Status::OK();
}
#ifdef __cplusplus
extern "C" {
@@ -26,10 +63,9 @@ extern "C" {
JNIEXPORT jbyteArray JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativePut__Lio_ray_runtime_object_NativeRayObject_2(
JNIEnv *env, jclass, jobject obj) {
auto ray_object = JavaNativeRayObjectToNativeRayObject(env, obj);
RAY_CHECK(ray_object != nullptr);
ray::ObjectID object_id;
auto status = ray::CoreWorkerProcess::GetCoreWorker().Put(*ray_object, {}, &object_id);
auto status = PutSerializedObject(env, obj, /*object_id=*/ray::ObjectID::Nil(),
/*out_object_id=*/&object_id, /*pin_object=*/true);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return IdToJavaByteArray<ray::ObjectID>(env, object_id);
}
@@ -38,9 +74,10 @@ JNIEXPORT void JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativePut___3BLio_ray_runtime_object_NativeRayObject_2(
JNIEnv *env, jclass, jbyteArray objectId, jobject obj) {
auto object_id = JavaByteArrayToId<ray::ObjectID>(env, objectId);
auto ray_object = JavaNativeRayObjectToNativeRayObject(env, obj);
RAY_CHECK(ray_object != nullptr);
auto status = ray::CoreWorkerProcess::GetCoreWorker().Put(*ray_object, {}, object_id);
ray::ObjectID dummy_object_id;
auto status =
PutSerializedObject(env, obj, object_id,
/*out_object_id=*/&dummy_object_id, /*pin_object=*/true);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
@@ -71,7 +108,10 @@ JNIEXPORT jobject JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeWai
object_ids, (int)numObjects, (int64_t)timeoutMs, &results);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return NativeVectorToJavaList<bool>(env, results, [](JNIEnv *env, const bool &item) {
return env->NewObject(java_boolean_class, java_boolean_init, (jboolean)item);
jobject java_item =
env->NewObject(java_boolean_class, java_boolean_init, (jboolean)item);
RAY_CHECK_JAVA_EXCEPTION(env);
return java_item;
});
}
@@ -88,6 +128,49 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeDelete
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
JNIEXPORT void JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativeAddLocalReference(
JNIEnv *env, jclass, jbyteArray workerId, jbyteArray objectId) {
auto worker_id = JavaByteArrayToId<ray::WorkerID>(env, workerId);
auto object_id = JavaByteArrayToId<ray::ObjectID>(env, objectId);
auto core_worker = ray::CoreWorkerProcess::TryGetWorker(worker_id);
RAY_CHECK(core_worker);
core_worker->AddLocalReference(object_id);
}
JNIEXPORT void JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativeRemoveLocalReference(
JNIEnv *env, jclass, jbyteArray workerId, jbyteArray objectId) {
auto worker_id = JavaByteArrayToId<ray::WorkerID>(env, workerId);
auto object_id = JavaByteArrayToId<ray::ObjectID>(env, objectId);
// We can't control the timing of Java GC, so it's normal that this method is called but
// core worker is shutting down (or already shut down). If we can't get a core worker
// instance here, skip calling the `RemoveLocalReference` method.
auto core_worker = ray::CoreWorkerProcess::TryGetWorker(worker_id);
if (core_worker) {
core_worker->RemoveLocalReference(object_id);
}
}
JNIEXPORT jobject JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativeGetAllReferenceCounts(JNIEnv *env,
jclass) {
auto reference_counts = ray::CoreWorkerProcess::GetCoreWorker().GetAllReferenceCounts();
return NativeMapToJavaMap<ray::ObjectID, std::pair<size_t, size_t>>(
env, reference_counts,
[](JNIEnv *env, const ray::ObjectID &key) {
return IdToJavaByteArray<ObjectID>(env, key);
},
[](JNIEnv *env, const std::pair<size_t, size_t> &value) {
jlongArray array = env->NewLongArray(2);
jlong *elements = env->GetLongArrayElements(array, nullptr);
elements[0] = static_cast<jlong>(value.first);
elements[1] = static_cast<jlong>(value.second);
env->ReleaseLongArrayElements(array, elements, 0);
return array;
});
}
#ifdef __cplusplus
}
#endif
@@ -65,6 +65,35 @@ JNIEXPORT jobject JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeWai
JNIEXPORT void JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeDelete(
JNIEnv *, jclass, jobject, jboolean, jboolean);
/*
* Class: io_ray_runtime_object_NativeObjectStore
* Method: nativeAddLocalReference
* Signature: ([B)V
*/
JNIEXPORT void JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativeAddLocalReference(JNIEnv *, jclass,
jbyteArray,
jbyteArray);
/*
* Class: io_ray_runtime_object_NativeObjectStore
* Method: nativeRemoveLocalReference
* Signature: ([B)V
*/
JNIEXPORT void JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativeRemoveLocalReference(JNIEnv *, jclass,
jbyteArray,
jbyteArray);
/*
* Class: io_ray_runtime_object_NativeObjectStore
* Method: nativeGetAllReferenceCounts
* Signature: ()Ljava/util/Map;
*/
JNIEXPORT jobject JNICALL
Java_io_ray_runtime_object_NativeObjectStore_nativeGetAllReferenceCounts(JNIEnv *,
jclass);
#ifdef __cplusplus
}
#endif
@@ -76,7 +76,6 @@ inline std::vector<std::unique_ptr<ray::TaskArg>> ToTaskArgs(JNIEnv *env, jobjec
inline std::unordered_map<std::string, double> ToResources(JNIEnv *env,
jobject java_resources) {
std::unordered_map<std::string, double> resources;
return JavaMapToNativeMap<std::string, double>(
env, java_resources,
[](JNIEnv *env, jobject java_key) {
+20
View File
@@ -34,6 +34,10 @@ jmethodID java_array_list_init_with_capacity;
jclass java_map_class;
jmethodID java_map_entry_set;
jmethodID java_map_put;
jclass java_hash_map_class;
jmethodID java_hash_map_init;
jclass java_set_class;
jmethodID java_set_iterator;
@@ -46,6 +50,9 @@ jclass java_map_entry_class;
jmethodID java_map_entry_get_key;
jmethodID java_map_entry_get_value;
jclass java_system_class;
jmethodID java_system_gc;
jclass java_ray_exception_class;
jclass java_jni_exception_util_class;
@@ -86,6 +93,7 @@ jclass java_native_ray_object_class;
jmethodID java_native_ray_object_init;
jfieldID java_native_ray_object_data;
jfieldID java_native_ray_object_metadata;
jfieldID java_native_ray_object_contained_object_ids;
jclass java_task_executor_class;
jmethodID java_task_executor_parse_function_arguments;
@@ -135,6 +143,11 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
java_map_class = LoadClass(env, "java/util/Map");
java_map_entry_set = env->GetMethodID(java_map_class, "entrySet", "()Ljava/util/Set;");
java_map_put = env->GetMethodID(
java_map_class, "put", "(Ljava/lang/Object;Ljava/lang/Object;)Ljava/lang/Object;");
java_hash_map_class = LoadClass(env, "java/util/HashMap");
java_hash_map_init = env->GetMethodID(java_hash_map_class, "<init>", "()V");
java_set_class = LoadClass(env, "java/util/Set");
java_set_iterator =
@@ -151,6 +164,9 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
java_map_entry_get_value =
env->GetMethodID(java_map_entry_class, "getValue", "()Ljava/lang/Object;");
java_system_class = LoadClass(env, "java/lang/System");
java_system_gc = env->GetStaticMethodID(java_system_class, "gc", "()V");
java_ray_exception_class = LoadClass(env, "io/ray/api/exception/RayException");
java_jni_exception_util_class = LoadClass(env, "io/ray/runtime/util/JniExceptionUtil");
@@ -220,6 +236,8 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
env->GetFieldID(java_native_ray_object_class, "data", "[B");
java_native_ray_object_metadata =
env->GetFieldID(java_native_ray_object_class, "metadata", "[B");
java_native_ray_object_contained_object_ids = env->GetFieldID(
java_native_ray_object_class, "containedObjectIds", "Ljava/util/List;");
java_task_executor_class = LoadClass(env, "io/ray/runtime/task/TaskExecutor");
java_task_executor_parse_function_arguments = env->GetMethodID(
@@ -241,9 +259,11 @@ void JNI_OnUnload(JavaVM *vm, void *reserved) {
env->DeleteGlobalRef(java_list_class);
env->DeleteGlobalRef(java_array_list_class);
env->DeleteGlobalRef(java_map_class);
env->DeleteGlobalRef(java_hash_map_class);
env->DeleteGlobalRef(java_set_class);
env->DeleteGlobalRef(java_iterator_class);
env->DeleteGlobalRef(java_map_entry_class);
env->DeleteGlobalRef(java_system_class);
env->DeleteGlobalRef(java_ray_exception_class);
env->DeleteGlobalRef(java_jni_exception_util_class);
env->DeleteGlobalRef(java_base_id_class);
+46 -3
View File
@@ -58,6 +58,13 @@ extern jmethodID java_array_list_init_with_capacity;
extern jclass java_map_class;
/// entrySet method of Map interface
extern jmethodID java_map_entry_set;
/// put method of Map interface
extern jmethodID java_map_put;
/// HashMap class
extern jclass java_hash_map_class;
/// Constructor of HashMap class
extern jmethodID java_hash_map_init;
/// Set interface
extern jclass java_set_class;
@@ -78,6 +85,11 @@ extern jmethodID java_map_entry_get_key;
/// getValue method of Map.Entry interface
extern jmethodID java_map_entry_get_value;
/// System class
extern jclass java_system_class;
/// gc method of System class
extern jmethodID java_system_gc;
/// RayException class
extern jclass java_ray_exception_class;
@@ -149,6 +161,8 @@ extern jmethodID java_native_ray_object_init;
extern jfieldID java_native_ray_object_data;
/// metadata field of NativeRayObject class
extern jfieldID java_native_ray_object_metadata;
// containedObjectIds field of NativeRayObject class
extern jfieldID java_native_ray_object_contained_object_ids;
/// TaskExecutor class
extern jclass java_task_executor_class;
@@ -323,6 +337,7 @@ inline jobject NativeVectorToJavaList(
jobject java_list =
env->NewObject(java_array_list_class, java_array_list_init_with_capacity,
(jint)native_vector.size());
RAY_CHECK_JAVA_EXCEPTION(env);
for (const auto &item : native_vector) {
auto element = element_converter(env, item);
env->CallVoidMethod(java_list, java_list_add, element);
@@ -364,10 +379,11 @@ inline std::unordered_map<key_type, value_type> JavaMapToNativeMap(
RAY_CHECK_JAVA_EXCEPTION(env);
jobject map_entry = env->CallObjectMethod(iterator, java_iterator_next);
RAY_CHECK_JAVA_EXCEPTION(env);
auto java_key = (jstring)env->CallObjectMethod(map_entry, java_map_entry_get_key);
auto java_key = env->CallObjectMethod(map_entry, java_map_entry_get_key);
RAY_CHECK_JAVA_EXCEPTION(env);
key_type key = key_converter(env, java_key);
auto java_value = env->CallObjectMethod(map_entry, java_map_entry_get_value);
RAY_CHECK_JAVA_EXCEPTION(env);
value_type value = value_converter(env, java_value);
native_map.emplace(key, value);
env->DeleteLocalRef(java_key);
@@ -381,6 +397,25 @@ inline std::unordered_map<key_type, value_type> JavaMapToNativeMap(
return native_map;
}
/// Convert a C++ std::unordered_map<?, ?> to a Java Map<?, ?>
template <typename key_type, typename value_type>
inline jobject NativeMapToJavaMap(
JNIEnv *env, const std::unordered_map<key_type, value_type> &native_map,
const std::function<jobject(JNIEnv *, const key_type &)> &key_converter,
const std::function<jobject(JNIEnv *, const value_type &)> &value_converter) {
jobject java_map = env->NewObject(java_hash_map_class, java_hash_map_init);
RAY_CHECK_JAVA_EXCEPTION(env);
for (const auto &entry : native_map) {
jobject java_key = key_converter(env, entry.first);
jobject java_value = value_converter(env, entry.second);
env->CallObjectMethod(java_map, java_map_put, java_key, java_value);
RAY_CHECK_JAVA_EXCEPTION(env);
env->DeleteLocalRef(java_key);
env->DeleteLocalRef(java_value);
}
return java_map;
}
/// Convert a C++ ray::Buffer to a Java byte array.
inline jbyteArray NativeBufferToJavaByteArray(JNIEnv *env,
const std::shared_ptr<ray::Buffer> buffer) {
@@ -423,9 +458,16 @@ inline std::shared_ptr<ray::RayObject> JavaNativeRayObjectToNativeRayObject(
if (metadata_buffer && metadata_buffer->Size() == 0) {
metadata_buffer = nullptr;
}
// TODO: Support nested IDs for Java.
auto java_contained_ids =
env->GetObjectField(java_obj, java_native_ray_object_contained_object_ids);
std::vector<ray::ObjectID> contained_object_ids;
JavaListToNativeVector<ray::ObjectID>(
env, java_contained_ids, &contained_object_ids, [](JNIEnv *env, jobject id) {
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
});
return std::make_shared<ray::RayObject>(data_buffer, metadata_buffer,
std::vector<ray::ObjectID>());
contained_object_ids);
}
/// Convert a C++ ray::RayObject to a Java NativeRayObject.
@@ -438,6 +480,7 @@ inline jobject NativeRayObjectToJavaNativeRayObject(
auto java_metadata = NativeBufferToJavaByteArray(env, rayObject->GetMetadata());
auto java_obj = env->NewObject(java_native_ray_object_class,
java_native_ray_object_init, java_data, java_metadata);
RAY_CHECK_JAVA_EXCEPTION(env);
env->DeleteLocalRef(java_metadata);
env->DeleteLocalRef(java_data);
return java_obj;