Skip to content

Commit 8c488d6

Browse files
Nguyen Ngoc LongNguyen Ngoc Long
authored andcommitted
fix-anr-when-checking-if-model-downloaded
1 parent dc355db commit 8c488d6

4 files changed

Lines changed: 256 additions & 168 deletions

File tree

Lines changed: 72 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,9 @@
11
package com.google_mlkit_commons;
22

3-
import com.google.android.gms.tasks.Task;
4-
import com.google.android.gms.tasks.Tasks;
53
import com.google.mlkit.common.model.DownloadConditions;
64
import com.google.mlkit.common.model.RemoteModel;
75
import com.google.mlkit.common.model.RemoteModelManager;
86

9-
import java.util.concurrent.Callable;
10-
import java.util.concurrent.ExecutionException;
11-
import java.util.concurrent.ExecutorService;
12-
import java.util.concurrent.Executors;
13-
import java.util.concurrent.Future;
14-
157
import io.flutter.plugin.common.MethodCall;
168
import io.flutter.plugin.common.MethodChannel;
179

@@ -20,16 +12,25 @@ public class GenericModelManager {
2012
private static final String DELETE = "delete";
2113
private static final String CHECK = "check";
2214

23-
public RemoteModelManager remoteModelManager = RemoteModelManager.getInstance();
15+
public interface CheckModelIsDownloadedCallback {
16+
void onModelDownloaded(Boolean isDownloaded);
17+
18+
void onError(Exception e);
19+
}
2420

25-
//To avoid downloading models in the main thread as they are around 20MB and may crash the app.
26-
private final ExecutorService executorService = Executors.newCachedThreadPool();
21+
public RemoteModelManager remoteModelManager = RemoteModelManager.getInstance();
2722

2823
public void manageModel(final RemoteModel model, final MethodCall call, final MethodChannel.Result result) {
2924
String task = call.argument("task");
25+
26+
if (task == null) {
27+
result.notImplemented();
28+
return;
29+
}
30+
3031
switch (task) {
3132
case DOWNLOAD:
32-
boolean isWifiReqRequired = call.argument("wifi");
33+
boolean isWifiReqRequired = Boolean.TRUE.equals(call.argument("wifi"));
3334
DownloadConditions downloadConditions;
3435
if (isWifiReqRequired)
3536
downloadConditions = new DownloadConditions.Builder().requireWifi().build();
@@ -41,52 +42,77 @@ public void manageModel(final RemoteModel model, final MethodCall call, final Me
4142
deleteModel(model, result);
4243
break;
4344
case CHECK:
44-
Boolean downloaded = isModelDownloaded(model);
45-
if (downloaded != null) result.success(downloaded);
46-
else result.error("error", null, null);
45+
isModelDownloaded(
46+
model,
47+
new CheckModelIsDownloadedCallback() {
48+
@Override
49+
public void onModelDownloaded(Boolean isDownloaded) {
50+
result.success(isDownloaded);
51+
}
52+
53+
@Override
54+
public void onError(Exception e) {
55+
result.error("error", e.toString(), null);
56+
}
57+
}
58+
);
4759
break;
4860
default:
4961
result.notImplemented();
5062
}
5163
}
5264

5365
public void downloadModel(RemoteModel remoteModel, DownloadConditions downloadConditions, final MethodChannel.Result result) {
54-
if (isModelDownloaded(remoteModel)) {
55-
result.success("success");
56-
return;
57-
}
58-
remoteModelManager.download(remoteModel, downloadConditions).addOnSuccessListener(aVoid -> result.success("success")).addOnFailureListener(e -> result.error("error", e.toString(), null));
59-
}
66+
isModelDownloaded(
67+
remoteModel,
68+
new CheckModelIsDownloadedCallback() {
69+
@Override
70+
public void onModelDownloaded(Boolean isDownloaded) {
71+
if (isDownloaded) {
72+
result.success("success");
73+
return;
74+
}
6075

61-
public void deleteModel(RemoteModel remoteModel, final MethodChannel.Result result) {
62-
if (!isModelDownloaded(remoteModel)) {
63-
result.success("success");
64-
return;
65-
}
66-
remoteModelManager.deleteDownloadedModel(remoteModel).addOnSuccessListener(aVoid -> result.success("success")).addOnFailureListener(e -> result.error("error", e.toString(), null));
67-
}
76+
remoteModelManager.download(remoteModel, downloadConditions)
77+
.addOnSuccessListener(aVoid -> result.success("success"))
78+
.addOnFailureListener(e -> result.error("error", e.toString(), null));
79+
}
6880

69-
public Boolean isModelDownloaded(RemoteModel model) {
70-
IsModelDownloaded myCallable = new IsModelDownloaded(remoteModelManager.isModelDownloaded(model));
71-
Future<Boolean> taskResult = executorService.submit(myCallable);
72-
try {
73-
return taskResult.get();
74-
} catch (InterruptedException | ExecutionException e) {
75-
e.printStackTrace();
76-
}
77-
return null;
81+
@Override
82+
public void onError(Exception e) {
83+
result.error("error", e.toString(), null);
84+
}
85+
}
86+
);
7887
}
79-
}
8088

81-
class IsModelDownloaded implements Callable<Boolean> {
82-
final Task<Boolean> booleanTask;
89+
public void deleteModel(RemoteModel remoteModel, final MethodChannel.Result result) {
90+
isModelDownloaded(remoteModel, new CheckModelIsDownloadedCallback() {
91+
@Override
92+
public void onModelDownloaded(Boolean isDownloaded) {
93+
if (!isDownloaded) {
94+
result.success("success");
95+
return;
96+
}
97+
remoteModelManager.deleteDownloadedModel(remoteModel)
98+
.addOnSuccessListener(aVoid -> result.success("success"))
99+
.addOnFailureListener(e -> result.error("error", e.toString(), null));
100+
}
83101

84-
public IsModelDownloaded(Task<Boolean> booleanTask) {
85-
this.booleanTask = booleanTask;
102+
@Override
103+
public void onError(Exception e) {
104+
result.error("error", e.toString(), null);
105+
}
106+
});
86107
}
87108

88-
@Override
89-
public Boolean call() throws Exception {
90-
return Tasks.await(booleanTask);
109+
public void isModelDownloaded(RemoteModel model, CheckModelIsDownloadedCallback callback) {
110+
try {
111+
remoteModelManager.isModelDownloaded(model)
112+
.addOnFailureListener(e -> callback.onError(e))
113+
.addOnSuccessListener(isDownloaded -> callback.onModelDownloaded(isDownloaded));
114+
} catch (Exception e) {
115+
callback.onError(e);
116+
}
91117
}
92-
}
118+
}

packages/google_mlkit_digital_ink_recognition/android/src/main/java/com/google_mlkit_digital_ink_recognition/DigitalInkRecognizer.java

Lines changed: 77 additions & 63 deletions
Original file line numberDiff line numberDiff line change
@@ -54,71 +54,85 @@ private void handleDetection(MethodCall call, final MethodChannel.Result result)
5454
DigitalInkRecognitionModel model = getModel(tag, result);
5555
if (model == null)
5656
return;
57-
if (!genericModelManager.isModelDownloaded(model)) {
58-
result.error("Model Error", "Model has not been downloaded yet ", null);
59-
return;
60-
}
61-
62-
String id = call.argument("id");
63-
com.google.mlkit.vision.digitalink.DigitalInkRecognizer recognizer = instances.get(id);
64-
if (recognizer == null) {
65-
recognizer = DigitalInkRecognition.getClient(DigitalInkRecognizerOptions.builder(model).build());
66-
instances.put(id, recognizer);
67-
}
6857

69-
Map<String, Object> inkMap = call.argument("ink");
70-
List<Map<String, Object>> strokeList = (List<Map<String, Object>>) inkMap.get("strokes");
71-
Ink.Builder inkBuilder = Ink.builder();
72-
for (final Map<String, Object> strokeMap : strokeList) {
73-
Ink.Stroke.Builder strokeBuilder = Ink.Stroke.builder();
74-
List<Map<String, Object>> pointsList = (List<Map<String, Object>>) strokeMap.get("points");
75-
for (final Map<String, Object> point : pointsList) {
76-
float x = (float) (double) point.get("x");
77-
float y = (float) (double) point.get("y");
78-
Object t0 = point.get("t");
79-
long t;
80-
if (t0 instanceof Integer) {
81-
t = (int) t0;
82-
} else {
83-
t = (long) t0;
58+
genericModelManager.isModelDownloaded(
59+
model,
60+
new GenericModelManager.CheckModelIsDownloadedCallback() {
61+
@Override
62+
public void onModelDownloaded(Boolean isDownloaded) {
63+
if (!isDownloaded) {
64+
result.error("Model Error", "Model has not been downloaded yet ", null);
65+
return;
66+
}
67+
68+
String id = call.argument("id");
69+
com.google.mlkit.vision.digitalink.DigitalInkRecognizer recognizer = instances.get(id);
70+
if (recognizer == null) {
71+
recognizer = DigitalInkRecognition.getClient(DigitalInkRecognizerOptions.builder(model).build());
72+
instances.put(id, recognizer);
73+
}
74+
75+
Map<String, Object> inkMap = call.argument("ink");
76+
List<Map<String, Object>> strokeList = (List<Map<String, Object>>) inkMap.get("strokes");
77+
Ink.Builder inkBuilder = Ink.builder();
78+
for (final Map<String, Object> strokeMap : strokeList) {
79+
Ink.Stroke.Builder strokeBuilder = Ink.Stroke.builder();
80+
List<Map<String, Object>> pointsList = (List<Map<String, Object>>) strokeMap.get("points");
81+
for (final Map<String, Object> point : pointsList) {
82+
float x = (float) (double) point.get("x");
83+
float y = (float) (double) point.get("y");
84+
Object t0 = point.get("t");
85+
long t;
86+
if (t0 instanceof Integer) {
87+
t = (int) t0;
88+
} else {
89+
t = (long) t0;
90+
}
91+
Ink.Point strokePoint = Ink.Point.create(x, y, t);
92+
strokeBuilder.addPoint(strokePoint);
93+
}
94+
inkBuilder.addStroke(strokeBuilder.build());
95+
}
96+
Ink ink = inkBuilder.build();
97+
98+
RecognitionContext context = null;
99+
Map<String, Object> contextMap = call.argument("context");
100+
if (contextMap != null) {
101+
RecognitionContext.Builder builder = RecognitionContext.builder();
102+
String preContext = (String) contextMap.get("preContext");
103+
if (preContext != null) {
104+
builder.setPreContext(preContext);
105+
} else {
106+
builder.setPreContext("");
107+
}
108+
109+
Map<String, Object> writingAreaMap = (Map<String, Object>) contextMap.get("writingArea");
110+
if (writingAreaMap != null) {
111+
float width = (float) (double) writingAreaMap.get("width");
112+
float height = (float) (double) writingAreaMap.get("height");
113+
builder.setWritingArea(new WritingArea(width, height));
114+
}
115+
116+
context = builder.build();
117+
}
118+
119+
if (context != null) {
120+
recognizer.recognize(ink, context)
121+
.addOnSuccessListener(recognitionResult -> process(recognitionResult, result))
122+
.addOnFailureListener(e -> result.error("recognition Error", e.toString(), null));
123+
} else {
124+
recognizer.recognize(ink)
125+
.addOnSuccessListener(recognitionResult -> process(recognitionResult, result))
126+
.addOnFailureListener(e -> result.error("recognition Error", e.toString(), null));
127+
}
128+
}
129+
130+
@Override
131+
public void onError(Exception e) {
132+
result.error("error", e.toString(), null);
133+
}
84134
}
85-
Ink.Point strokePoint = Ink.Point.create(x, y, t);
86-
strokeBuilder.addPoint(strokePoint);
87-
}
88-
inkBuilder.addStroke(strokeBuilder.build());
89-
}
90-
Ink ink = inkBuilder.build();
91-
92-
RecognitionContext context = null;
93-
Map<String, Object> contextMap = call.argument("context");
94-
if (contextMap != null) {
95-
RecognitionContext.Builder builder = RecognitionContext.builder();
96-
String preContext = (String) contextMap.get("preContext");
97-
if (preContext != null) {
98-
builder.setPreContext(preContext);
99-
} else {
100-
builder.setPreContext("");
101-
}
102-
103-
Map<String, Object> writingAreaMap = (Map<String, Object>) contextMap.get("writingArea");
104-
if (writingAreaMap != null) {
105-
float width = (float) (double) writingAreaMap.get("width");
106-
float height = (float) (double) writingAreaMap.get("height");
107-
builder.setWritingArea(new WritingArea(width, height));
108-
}
109-
110-
context = builder.build();
111-
}
112-
113-
if (context != null) {
114-
recognizer.recognize(ink, context)
115-
.addOnSuccessListener(recognitionResult -> process(recognitionResult, result))
116-
.addOnFailureListener(e -> result.error("recognition Error", e.toString(), null));
117-
} else {
118-
recognizer.recognize(ink)
119-
.addOnSuccessListener(recognitionResult -> process(recognitionResult, result))
120-
.addOnFailureListener(e -> result.error("recognition Error", e.toString(), null));
121-
}
135+
);
122136
}
123137

124138
private void process(RecognitionResult recognitionResult, final MethodChannel.Result result) {

packages/google_mlkit_image_labeling/android/src/main/java/com/google_mlkit_image_labeling/ImageLabelDetector.java

Lines changed: 47 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -79,12 +79,53 @@ private void handleDetection(MethodCall call, final MethodChannel.Result result)
7979
CustomImageLabelerOptions labelerOptions = getLocalOptions(options);
8080
imageLabeler = ImageLabeling.getClient(labelerOptions);
8181
} else if (type.equals("remote")) {
82-
CustomImageLabelerOptions labelerOptions = getRemoteOptions(options);
83-
if (labelerOptions == null) {
84-
result.error("Error Model has not been downloaded yet", "Model has not been downloaded yet", "Model has not been downloaded yet");
85-
return;
86-
}
87-
imageLabeler = ImageLabeling.getClient(labelerOptions);
82+
float confidenceThreshold = (float) (double) options.get("confidenceThreshold");
83+
int maxCount = (int) options.get("maxCount");
84+
String name = (String) options.get("modelName");
85+
86+
FirebaseModelSource firebaseModelSource = new FirebaseModelSource.Builder(name).build();
87+
CustomRemoteModel remoteModel = new CustomRemoteModel.Builder(firebaseModelSource).build();
88+
89+
genericModelManager.isModelDownloaded(
90+
remoteModel,
91+
new GenericModelManager.CheckModelIsDownloadedCallback() {
92+
@Override
93+
public void onModelDownloaded(Boolean isDownloaded) {
94+
if (!isDownloaded) {
95+
result.error("Error Model has not been downloaded yet", "Model has not been downloaded yet", "Model has not been downloaded yet");
96+
return;
97+
}
98+
99+
CustomImageLabelerOptions labelerOptions = new CustomImageLabelerOptions.Builder(remoteModel)
100+
.setConfidenceThreshold(confidenceThreshold)
101+
.setMaxResultCount(maxCount)
102+
.build();
103+
104+
ImageLabeling.getClient(labelerOptions).process(inputImage)
105+
.addOnSuccessListener(imageLabels -> {
106+
List<Map<String, Object>> labels = new ArrayList<>(imageLabels.size());
107+
for (ImageLabel label : imageLabels) {
108+
Map<String, Object> labelData = new HashMap<>();
109+
labelData.put("text", label.getText());
110+
labelData.put("confidence", label.getConfidence());
111+
labelData.put("index", label.getIndex());
112+
labels.add(labelData);
113+
}
114+
115+
result.success(labels);
116+
})
117+
.addOnFailureListener(e -> result.error("ImageLabelDetectorError", e.toString(), null));
118+
;
119+
}
120+
121+
@Override
122+
public void onError(Exception e) {
123+
result.error("Error", e.getMessage(), e);
124+
}
125+
}
126+
);
127+
128+
return;
88129
} else {
89130
String error = "Invalid model type: " + type;
90131
result.error(type, error, error);
@@ -131,24 +172,6 @@ private CustomImageLabelerOptions getLocalOptions(Map<String, Object> labelerOpt
131172
.build();
132173
}
133174

134-
//Options for labeler to work with custom model.
135-
private CustomImageLabelerOptions getRemoteOptions(Map<String, Object> labelerOptions) {
136-
float confidenceThreshold = (float) (double) labelerOptions.get("confidenceThreshold");
137-
int maxCount = (int) labelerOptions.get("maxCount");
138-
String name = (String) labelerOptions.get("modelName");
139-
140-
FirebaseModelSource firebaseModelSource = new FirebaseModelSource.Builder(name).build();
141-
CustomRemoteModel remoteModel = new CustomRemoteModel.Builder(firebaseModelSource).build();
142-
if (!genericModelManager.isModelDownloaded(remoteModel)) {
143-
return null;
144-
}
145-
146-
return new CustomImageLabelerOptions.Builder(remoteModel)
147-
.setConfidenceThreshold(confidenceThreshold)
148-
.setMaxResultCount(maxCount)
149-
.build();
150-
}
151-
152175
private void closeDetector(MethodCall call) {
153176
String id = call.argument("id");
154177
ImageLabeler imageLabeler = instances.get(id);

0 commit comments

Comments
 (0)