From 253f13ad622139ca41b134fc1611f6d2907ca0a9 Mon Sep 17 00:00:00 2001 From: Sebastian Schmidt Date: Fri, 28 Apr 2023 20:31:57 -0700 Subject: [PATCH] Invoke callback for InteractiveSegmenter while C++ Packets are active PiperOrigin-RevId: 528053621 --- .../interactive_segmenter.ts | 24 +++++++++++++-- .../interactive_segmenter_test.ts | 30 +++++++++++++++++++ 2 files changed, 51 insertions(+), 3 deletions(-) diff --git a/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter.ts b/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter.ts index f127792c..67d6ec3f 100644 --- a/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter.ts +++ b/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter.ts @@ -86,6 +86,7 @@ export class InteractiveSegmenter extends VisionTaskRunner { private result: InteractiveSegmenterResult = {}; private outputCategoryMask = DEFAULT_OUTPUT_CATEGORY_MASK; private outputConfidenceMasks = DEFAULT_OUTPUT_CONFIDENCE_MASKS; + private userCallback: InteractiveSegmenterCallback = () => {}; private readonly options: ImageSegmenterGraphOptionsProto; private readonly segmenterOptions: SegmenterOptionsProto; @@ -241,21 +242,32 @@ export class InteractiveSegmenter extends VisionTaskRunner { typeof imageProcessingOptionsOrCallback !== 'function' ? imageProcessingOptionsOrCallback : {}; - const userCallback = - typeof imageProcessingOptionsOrCallback === 'function' ? + this.userCallback = typeof imageProcessingOptionsOrCallback === 'function' ? imageProcessingOptionsOrCallback : callback!; this.reset(); this.processRenderData(roi, this.getSynctheticTimestamp()); this.processImageData(image, imageProcessingOptions); - userCallback(this.result); + this.userCallback = () => {}; } private reset(): void { this.result = {}; } + /** Invokes the user callback once all data has been received. */ + private maybeInvokeCallback(): void { + if (this.outputConfidenceMasks && !('confidenceMasks' in this.result)) { + return; + } + if (this.outputCategoryMask && !('categoryMask' in this.result)) { + return; + } + + this.userCallback(this.result); + } + /** Updates the MediaPipe graph configuration. */ protected override refreshGraph(): void { const graphConfig = new CalculatorGraphConfig(); @@ -286,10 +298,13 @@ export class InteractiveSegmenter extends VisionTaskRunner { this.result.confidenceMasks = masks.map(wasmImage => this.convertToMPImage(wasmImage)); this.setLatestOutputTimestamp(timestamp); + this.maybeInvokeCallback(); }); this.graphRunner.attachEmptyPacketListener( CONFIDENCE_MASKS_STREAM, timestamp => { + this.result.confidenceMasks = undefined; this.setLatestOutputTimestamp(timestamp); + this.maybeInvokeCallback(); }); } @@ -301,10 +316,13 @@ export class InteractiveSegmenter extends VisionTaskRunner { CATEGORY_MASK_STREAM, (mask, timestamp) => { this.result.categoryMask = this.convertToMPImage(mask); this.setLatestOutputTimestamp(timestamp); + this.maybeInvokeCallback(); }); this.graphRunner.attachEmptyPacketListener( CATEGORY_MASK_STREAM, timestamp => { + this.result.categoryMask = undefined; this.setLatestOutputTimestamp(timestamp); + this.maybeInvokeCallback(); }); } diff --git a/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter_test.ts b/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter_test.ts index cbaf7663..84ecde00 100644 --- a/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter_test.ts +++ b/mediapipe/tasks/web/vision/interactive_segmenter/interactive_segmenter_test.ts @@ -252,4 +252,34 @@ describe('InteractiveSegmenter', () => { }); }); }); + + it('invokes listener once masks are avaiblae', async () => { + const categoryMask = new Uint8ClampedArray([1]); + const confidenceMask = new Float32Array([0.0]); + let listenerCalled = false; + + await interactiveSegmenter.setOptions( + {outputCategoryMask: true, outputConfidenceMasks: true}); + + // Pass the test data to our listener + interactiveSegmenter.fakeWasmModule._waitUntilIdle.and.callFake(() => { + expect(listenerCalled).toBeFalse(); + interactiveSegmenter.categoryMaskListener! + ({data: categoryMask, width: 1, height: 1}, 1337); + expect(listenerCalled).toBeFalse(); + interactiveSegmenter.confidenceMasksListener!( + [ + {data: confidenceMask, width: 1, height: 1}, + ], + 1337); + expect(listenerCalled).toBeTrue(); + }); + + return new Promise(resolve => { + interactiveSegmenter.segment({} as HTMLImageElement, ROI, () => { + listenerCalled = true; + resolve(); + }); + }); + }); });