Merge "Load back-gesture model on background thread." into tm-qpr-dev

This commit is contained in:
Cutter Coryell
2022-08-04 16:32:53 +00:00
committed by Android (Google) Code Review

View File

@@ -60,6 +60,7 @@ import com.android.internal.policy.GestureNavigationSettingsObserver;
import com.android.internal.util.LatencyTracker; import com.android.internal.util.LatencyTracker;
import com.android.systemui.R; import com.android.systemui.R;
import com.android.systemui.broadcast.BroadcastDispatcher; import com.android.systemui.broadcast.BroadcastDispatcher;
import com.android.systemui.dagger.qualifiers.Background;
import com.android.systemui.dagger.qualifiers.Main; import com.android.systemui.dagger.qualifiers.Main;
import com.android.systemui.flags.FeatureFlags; import com.android.systemui.flags.FeatureFlags;
import com.android.systemui.flags.Flags; import com.android.systemui.flags.Flags;
@@ -82,6 +83,7 @@ import com.android.systemui.shared.tracing.ProtoTraceable;
import com.android.systemui.tracing.ProtoTracer; import com.android.systemui.tracing.ProtoTracer;
import com.android.systemui.tracing.nano.EdgeBackGestureHandlerProto; import com.android.systemui.tracing.nano.EdgeBackGestureHandlerProto;
import com.android.systemui.tracing.nano.SystemUiTraceProto; import com.android.systemui.tracing.nano.SystemUiTraceProto;
import com.android.systemui.util.Assert;
import com.android.wm.shell.back.BackAnimation; import com.android.wm.shell.back.BackAnimation;
import com.android.wm.shell.pip.Pip; import com.android.wm.shell.pip.Pip;
@@ -191,6 +193,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
private final int mDisplayId; private final int mDisplayId;
private final Executor mMainExecutor; private final Executor mMainExecutor;
private final Executor mBackgroundExecutor;
private final Rect mPipExcludedBounds = new Rect(); private final Rect mPipExcludedBounds = new Rect();
private final Rect mNavBarOverlayExcludedBounds = new Rect(); private final Rect mNavBarOverlayExcludedBounds = new Rect();
@@ -251,6 +254,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
private BackGestureTfClassifierProvider mBackGestureTfClassifierProvider; private BackGestureTfClassifierProvider mBackGestureTfClassifierProvider;
private Map<String, Integer> mVocab; private Map<String, Integer> mVocab;
private boolean mUseMLModel; private boolean mUseMLModel;
private boolean mMLModelIsLoading;
// minimum width below which we do not run the model // minimum width below which we do not run the model
private int mMLEnableWidth; private int mMLEnableWidth;
private float mMLModelThreshold; private float mMLModelThreshold;
@@ -318,6 +322,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
SysUiState sysUiState, SysUiState sysUiState,
PluginManager pluginManager, PluginManager pluginManager,
@Main Executor executor, @Main Executor executor,
@Background Executor backgroundExecutor,
BroadcastDispatcher broadcastDispatcher, BroadcastDispatcher broadcastDispatcher,
ProtoTracer protoTracer, ProtoTracer protoTracer,
NavigationModeController navigationModeController, NavigationModeController navigationModeController,
@@ -334,6 +339,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
mContext = context; mContext = context;
mDisplayId = context.getDisplayId(); mDisplayId = context.getDisplayId();
mMainExecutor = executor; mMainExecutor = executor;
mBackgroundExecutor = backgroundExecutor;
mOverviewProxyService = overviewProxyService; mOverviewProxyService = overviewProxyService;
mSysUiState = sysUiState; mSysUiState = sysUiState;
mPluginManager = pluginManager; mPluginManager = pluginManager;
@@ -631,28 +637,63 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
return; return;
} }
if (newState) { mUseMLModel = newState;
mBackGestureTfClassifierProvider = mBackGestureTfClassifierProviderProvider.get();
mMLModelThreshold = DeviceConfig.getFloat(DeviceConfig.NAMESPACE_SYSTEMUI, if (mUseMLModel) {
SystemUiDeviceConfigFlags.BACK_GESTURE_ML_MODEL_THRESHOLD, 0.9f); Assert.isMainThread();
if (mBackGestureTfClassifierProvider.isActive()) { if (mMLModelIsLoading) {
Trace.beginSection("EdgeBackGestureHandler#loadVocab"); Log.d(TAG, "Model tried to load while already loading.");
mVocab = mBackGestureTfClassifierProvider.loadVocab(mContext.getAssets());
Trace.endSection();
mUseMLModel = true;
return; return;
} }
} mMLModelIsLoading = true;
mBackgroundExecutor.execute(() -> loadMLModel());
mUseMLModel = false; } else if (mBackGestureTfClassifierProvider != null) {
if (mBackGestureTfClassifierProvider != null) {
mBackGestureTfClassifierProvider.release(); mBackGestureTfClassifierProvider.release();
mBackGestureTfClassifierProvider = null; mBackGestureTfClassifierProvider = null;
mVocab = null;
} }
} }
private void loadMLModel() {
BackGestureTfClassifierProvider provider = mBackGestureTfClassifierProviderProvider.get();
float threshold = DeviceConfig.getFloat(DeviceConfig.NAMESPACE_SYSTEMUI,
SystemUiDeviceConfigFlags.BACK_GESTURE_ML_MODEL_THRESHOLD, 0.9f);
Map<String, Integer> vocab = null;
if (provider != null && !provider.isActive()) {
provider.release();
provider = null;
Log.w(TAG, "Cannot load model because it isn't active");
}
if (provider != null) {
Trace.beginSection("EdgeBackGestureHandler#loadVocab");
vocab = provider.loadVocab(mContext.getAssets());
Trace.endSection();
}
BackGestureTfClassifierProvider finalProvider = provider;
Map<String, Integer> finalVocab = vocab;
mMainExecutor.execute(() -> onMLModelLoadFinished(finalProvider, finalVocab, threshold));
}
private void onMLModelLoadFinished(BackGestureTfClassifierProvider provider,
Map<String, Integer> vocab, float threshold) {
Assert.isMainThread();
mMLModelIsLoading = false;
if (!mUseMLModel) {
// This can happen if the user disables Gesture Nav while the model is loading.
if (provider != null) {
provider.release();
}
Log.d(TAG, "Model finished loading but isn't needed.");
return;
}
mBackGestureTfClassifierProvider = provider;
mVocab = vocab;
mMLModelThreshold = threshold;
}
private int getBackGesturePredictionsCategory(int x, int y, int app) { private int getBackGesturePredictionsCategory(int x, int y, int app) {
if (app == -1) { BackGestureTfClassifierProvider provider = mBackGestureTfClassifierProvider;
if (provider == null || app == -1) {
return -1; return -1;
} }
int distanceFromEdge; int distanceFromEdge;
@@ -673,7 +714,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
new long[]{(long) y}, new long[]{(long) y},
}; };
mMLResults = mBackGestureTfClassifierProvider.predict(featuresVector); mMLResults = provider.predict(featuresVector);
if (mMLResults == -1) { if (mMLResults == -1) {
return -1; return -1;
} }
@@ -1031,6 +1072,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
private final SysUiState mSysUiState; private final SysUiState mSysUiState;
private final PluginManager mPluginManager; private final PluginManager mPluginManager;
private final Executor mExecutor; private final Executor mExecutor;
private final Executor mBackgroundExecutor;
private final BroadcastDispatcher mBroadcastDispatcher; private final BroadcastDispatcher mBroadcastDispatcher;
private final ProtoTracer mProtoTracer; private final ProtoTracer mProtoTracer;
private final NavigationModeController mNavigationModeController; private final NavigationModeController mNavigationModeController;
@@ -1050,6 +1092,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
SysUiState sysUiState, SysUiState sysUiState,
PluginManager pluginManager, PluginManager pluginManager,
@Main Executor executor, @Main Executor executor,
@Background Executor backgroundExecutor,
BroadcastDispatcher broadcastDispatcher, BroadcastDispatcher broadcastDispatcher,
ProtoTracer protoTracer, ProtoTracer protoTracer,
NavigationModeController navigationModeController, NavigationModeController navigationModeController,
@@ -1067,6 +1110,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
mSysUiState = sysUiState; mSysUiState = sysUiState;
mPluginManager = pluginManager; mPluginManager = pluginManager;
mExecutor = executor; mExecutor = executor;
mBackgroundExecutor = backgroundExecutor;
mBroadcastDispatcher = broadcastDispatcher; mBroadcastDispatcher = broadcastDispatcher;
mProtoTracer = protoTracer; mProtoTracer = protoTracer;
mNavigationModeController = navigationModeController; mNavigationModeController = navigationModeController;
@@ -1089,6 +1133,7 @@ public class EdgeBackGestureHandler extends CurrentUserTracker
mSysUiState, mSysUiState,
mPluginManager, mPluginManager,
mExecutor, mExecutor,
mBackgroundExecutor,
mBroadcastDispatcher, mBroadcastDispatcher,
mProtoTracer, mProtoTracer,
mNavigationModeController, mNavigationModeController,