From ac8444f299d3f2a8eb9fdf9a921a8d9f2c651be1 Mon Sep 17 00:00:00 2001 From: Jhen-Jie Hong Date: Fri, 9 Jun 2023 07:17:33 +0800 Subject: [PATCH] [js/rn] Implement dispose native method (#16131) ### Description Implement `dispose` react native method. ### Motivation and Context Currently we are not able to release the memory used by model in JS runtime if we don't want to use it anymore, we can do that only by reload app on debug or restart app on release. --- .../reactnative/OnnxruntimeModuleTest.java | 3 + .../reactnative/OnnxruntimeModule.java | 55 ++++++++++++++++++- js/react_native/ios/OnnxruntimeModule.h | 2 + js/react_native/ios/OnnxruntimeModule.mm | 43 +++++++++++++-- .../OnnxruntimeModuleTest.mm | 6 ++ js/react_native/lib/backend.ts | 2 +- js/react_native/lib/binding.ts | 1 + 7 files changed, 105 insertions(+), 7 deletions(-) diff --git a/js/react_native/android/src/androidTest/java/ai/onnxruntime/reactnative/OnnxruntimeModuleTest.java b/js/react_native/android/src/androidTest/java/ai/onnxruntime/reactnative/OnnxruntimeModuleTest.java index ea5a7d2745..112d1c9860 100644 --- a/js/react_native/android/src/androidTest/java/ai/onnxruntime/reactnative/OnnxruntimeModuleTest.java +++ b/js/react_native/android/src/androidTest/java/ai/onnxruntime/reactnative/OnnxruntimeModuleTest.java @@ -135,6 +135,9 @@ public class OnnxruntimeModuleTest { Assert.fail(e.getMessage()); } } + + // test dispose + ortModule.dispose(sessionKey); } finally { mockSession.finishMocking(); } diff --git a/js/react_native/android/src/main/java/ai/onnxruntime/reactnative/OnnxruntimeModule.java b/js/react_native/android/src/main/java/ai/onnxruntime/reactnative/OnnxruntimeModule.java index 81b19e3829..685c8d0643 100644 --- a/js/react_native/android/src/main/java/ai/onnxruntime/reactnative/OnnxruntimeModule.java +++ b/js/react_native/android/src/main/java/ai/onnxruntime/reactnative/OnnxruntimeModule.java @@ -18,6 +18,7 @@ import android.util.Log; import androidx.annotation.NonNull; import androidx.annotation.RequiresApi; import com.facebook.react.bridge.Arguments; +import com.facebook.react.bridge.LifecycleEventListener; import com.facebook.react.bridge.Promise; import com.facebook.react.bridge.ReactApplicationContext; import com.facebook.react.bridge.ReactContextBaseJavaModule; @@ -42,7 +43,7 @@ import java.util.stream.Collectors; import java.util.stream.Stream; @RequiresApi(api = Build.VERSION_CODES.N) -public class OnnxruntimeModule extends ReactContextBaseJavaModule { +public class OnnxruntimeModule extends ReactContextBaseJavaModule implements LifecycleEventListener { private static ReactApplicationContext reactContext; private static OrtEnvironment ortEnvironment = OrtEnvironment.getEnvironment(); @@ -105,6 +106,22 @@ public class OnnxruntimeModule extends ReactContextBaseJavaModule { } } + /** + * React native binding API to dispose a session. + * + * @param key session key representing a session given at loadModel() + * @param promise output returning back to react native js + */ + @ReactMethod + public void dispose(String key, Promise promise) { + try { + dispose(key); + promise.resolve(null); + } catch (OrtException e) { + promise.reject("Failed to dispose session: " + e.getMessage(), e); + } + } + /** * React native binding API to run a model using given uri. * @@ -195,6 +212,19 @@ public class OnnxruntimeModule extends ReactContextBaseJavaModule { return resultMap; } + /** + * Dispose a model using given key. + * + * @param key a session key representing the session given at loadModel() + */ + public void dispose(String key) throws OrtException { + OrtSession ortSession = sessionMap.get(key); + if (ortSession != null) { + ortSession.close(); + sessionMap.remove(key); + } + } + /** * Run a model using given uri. * @@ -215,6 +245,7 @@ public class OnnxruntimeModule extends ReactContextBaseJavaModule { long startTime = System.currentTimeMillis(); Map feed = new HashMap<>(); Iterator iterator = ortSession.getInputNames().iterator(); + Result result = null; try { while (iterator.hasNext()) { String inputName = iterator.next(); @@ -252,7 +283,6 @@ public class OnnxruntimeModule extends ReactContextBaseJavaModule { Log.d("Duration", "createInputTensor: " + duration); startTime = System.currentTimeMillis(); - Result result = null; if (requestedOutputs != null) { result = ortSession.run(feed, requestedOutputs, runOptions); } else { @@ -270,6 +300,9 @@ public class OnnxruntimeModule extends ReactContextBaseJavaModule { } finally { OnnxValue.close(feed); + if (result != null) { + result.close(); + } } } @@ -358,4 +391,22 @@ public class OnnxruntimeModule extends ReactContextBaseJavaModule { return runOptions; } + + @Override + public void onHostResume() {} + + @Override + public void onHostPause() {} + + @Override + public void onHostDestroy() { + for (String key : sessionMap.keySet()) { + try { + dispose(key); + } catch (Exception e) { + Log.e("onHostDestroy", "Failed to dispose session: " + key, e); + } + } + sessionMap.clear(); + } } diff --git a/js/react_native/ios/OnnxruntimeModule.h b/js/react_native/ios/OnnxruntimeModule.h index 51d2e8fa78..4ea3d21f20 100644 --- a/js/react_native/ios/OnnxruntimeModule.h +++ b/js/react_native/ios/OnnxruntimeModule.h @@ -14,6 +14,8 @@ -(NSDictionary*)loadModelFromBuffer:(NSData*)modelData options:(NSDictionary*)options; +-(void)dispose:(NSString*)key; + -(NSDictionary*)run:(NSString*)url input:(NSDictionary*)input output:(NSArray*)output diff --git a/js/react_native/ios/OnnxruntimeModule.mm b/js/react_native/ios/OnnxruntimeModule.mm index 15ede97958..fa34583cb0 100644 --- a/js/react_native/ios/OnnxruntimeModule.mm +++ b/js/react_native/ios/OnnxruntimeModule.mm @@ -90,6 +90,25 @@ RCT_EXPORT_METHOD(loadModelFromBase64EncodedBuffer } } +/** + * React native binding API to dispose a session using given key from loadModel() + * + * @param key a model path location given at loadModel() + * @param resolve callback for returning output back to react native js + * @param reject callback for returning an error back to react native js + */ +RCT_EXPORT_METHOD(dispose + : (NSString *)key resolver + : (RCTPromiseResolveBlock)resolve rejecter + : (RCTPromiseRejectBlock)reject) { + @try { + [self dispose:key]; + resolve(nil); + } @catch (...) { + reject(@"onnxruntime", @"failed to dispose session", nil); + } +} + /** * React native binding API to run a model using given uri. * @@ -196,6 +215,25 @@ RCT_EXPORT_METHOD(run return resultMap; } +/** + * Dispose a session given a key. + * + * @param key a session key returned from loadModel() + */ +- (void)dispose:(NSString *)key { + NSValue *value = [sessionMap objectForKey:key]; + if (value == nil) { + NSException *exception = [NSException exceptionWithName:@"onnxruntime" + reason:@"can't find onnxruntime session" + userInfo:nil]; + @throw exception; + } + [sessionMap removeObjectForKey:key]; + SessionInfo *sessionInfo = (SessionInfo *)[value pointerValue]; + delete sessionInfo; + sessionInfo = nullptr; +} + /** * Run a model using given uri. * @@ -338,10 +376,7 @@ static NSDictionary *executionModeTable = @{@"sequential" : @(ORT_SEQUENTIAL), @ - (void)dealloc { NSEnumerator *iterator = [sessionMap keyEnumerator]; while (NSString *key = [iterator nextObject]) { - NSValue *value = [sessionMap objectForKey:key]; - SessionInfo *sessionInfo = (SessionInfo *)[value pointerValue]; - delete sessionInfo; - sessionInfo = nullptr; + [self dispose:key]; } } diff --git a/js/react_native/ios/OnnxruntimeModuleTest/OnnxruntimeModuleTest.mm b/js/react_native/ios/OnnxruntimeModuleTest/OnnxruntimeModuleTest.mm index 71fa0a44b7..86bf35229d 100644 --- a/js/react_native/ios/OnnxruntimeModuleTest/OnnxruntimeModuleTest.mm +++ b/js/react_native/ios/OnnxruntimeModuleTest/OnnxruntimeModuleTest.mm @@ -87,6 +87,12 @@ XCTAssertTrue([[resultMap objectForKey:@"output"] isEqualToDictionary:inputTensorMap]); XCTAssertTrue([[resultMap2 objectForKey:@"output"] isEqualToDictionary:inputTensorMap]); } + + // test dispose + { + [onnxruntimeModule dispose:sessionKey]; + [onnxruntimeModule dispose:sessionKey2]; + } } @end diff --git a/js/react_native/lib/backend.ts b/js/react_native/lib/backend.ts index 6fd7d7bea4..fa6d5ceaef 100644 --- a/js/react_native/lib/backend.ts +++ b/js/react_native/lib/backend.ts @@ -85,7 +85,7 @@ class OnnxruntimeSessionHandler implements SessionHandler { } async dispose(): Promise { - return Promise.resolve(); + return this.#inferenceSession.dispose(this.#key); } startProfiling(): void { diff --git a/js/react_native/lib/binding.ts b/js/react_native/lib/binding.ts index 4a4227e504..6d621fa3dc 100644 --- a/js/react_native/lib/binding.ts +++ b/js/react_native/lib/binding.ts @@ -65,6 +65,7 @@ export declare namespace Binding { interface InferenceSession { loadModel(modelPath: string, options: SessionOptions): Promise; loadModelFromBase64EncodedBuffer?(buffer: string, options: SessionOptions): Promise; + dispose(key: string): Promise; run(key: string, feeds: FeedsType, fetches: FetchesType, options: RunOptions): Promise; } }