diff --git a/mediapipe/tasks/web/vision/core/BUILD b/mediapipe/tasks/web/vision/core/BUILD index daeef060..c53247ba 100644 --- a/mediapipe/tasks/web/vision/core/BUILD +++ b/mediapipe/tasks/web/vision/core/BUILD @@ -91,6 +91,7 @@ mediapipe_ts_library( mediapipe_ts_library( name = "render_utils", srcs = ["render_utils.ts"], + deps = [":image"], ) jasmine_node_test( diff --git a/mediapipe/tasks/web/vision/core/render_utils.ts b/mediapipe/tasks/web/vision/core/render_utils.ts index 13af6176..892cd864 100644 --- a/mediapipe/tasks/web/vision/core/render_utils.ts +++ b/mediapipe/tasks/web/vision/core/render_utils.ts @@ -16,9 +16,11 @@ * limitations under the License. */ +import {MPImageChannelConverter} from '../../../../tasks/web/vision/core/image'; + // Pre-baked color table for a maximum of 12 classes. const CM_ALPHA = 128; -const COLOR_MAP = [ +const COLOR_MAP: Array<[number, number, number, number]> = [ [0, 0, 0, CM_ALPHA], // class 0 is BG = transparent [255, 0, 0, CM_ALPHA], // class 1 is red [0, 255, 0, CM_ALPHA], // class 2 is light green @@ -74,3 +76,9 @@ export function drawCategoryMask( } ctx.putImageData(new ImageData(rgbaArray, width, height), 0, 0); } + +/** The color converter we use in our demos. */ +export const RENDER_UTIL_CONVERTER: MPImageChannelConverter = { + floatToRGBAConverter: v => [128, 0, 0, v * 255], + uint8ToRGBAConverter: v => COLOR_MAP[v % COLOR_MAP.length], +}; diff --git a/mediapipe/tasks/web/vision/core/vision_task_runner.ts b/mediapipe/tasks/web/vision/core/vision_task_runner.ts index 5099d296..285dbf90 100644 --- a/mediapipe/tasks/web/vision/core/vision_task_runner.ts +++ b/mediapipe/tasks/web/vision/core/vision_task_runner.ts @@ -231,39 +231,41 @@ export abstract class VisionTaskRunner extends TaskRunner { */ protected convertToMPImage(wasmImage: WasmImage): MPImage { const {data, width, height} = wasmImage; + const pixels = width * height; + let container: ImageData|WebGLTexture|Uint8ClampedArray; if (data instanceof Uint8ClampedArray) { - let rgba: Uint8ClampedArray; - if (data.length === width * height * 4) { - rgba = data; - } else if (data.length === width * height * 3) { + if (data.length === pixels) { + container = data; // Mask + } else if (data.length === pixels * 3) { // TODO: Convert in C++ - rgba = new Uint8ClampedArray(width * height * 4); - for (let i = 0; i < width * height; ++i) { + const rgba = new Uint8ClampedArray(pixels * 4); + for (let i = 0; i < pixels; ++i) { rgba[4 * i] = data[3 * i]; rgba[4 * i + 1] = data[3 * i + 1]; rgba[4 * i + 2] = data[3 * i + 2]; rgba[4 * i + 3] = 255; } + container = new ImageData(rgba, width, height); + } else if (data.length ===pixels * 4) { + container = new ImageData(data, width, height); } else { - throw new Error( - `Unsupported channel count: ${data.length / width / height}`); + throw new Error(`Unsupported channel count: ${data.length/pixels}`); } - - return new MPImage( - [new ImageData(rgba, width, height)], - /* ownsImageBitmap= */ false, /* ownsWebGLTexture= */ false, - this.graphRunner.wasmModule.canvas!, this.shaderContext, width, - height); - } else if (data instanceof WebGLTexture) { - return new MPImage( - [data], /* ownsImageBitmap= */ false, /* ownsWebGLTexture= */ false, - this.graphRunner.wasmModule.canvas!, this.shaderContext, width, - height); - } else { - throw new Error( - `Cannot convert type ${data.constructor.name} to MPImage.`); + } else if (data instanceof Float32Array) { + if (data.length === pixels) { + container = data; // Mask + } else { + throw new Error(`Unsupported channel count: ${data.length/pixels}`); + } + } else { // WebGLTexture + container = data; } + + return new MPImage( + [container], /* ownsImageBitmap= */ false, /* ownsWebGLTexture= */ false, + this.graphRunner.wasmModule.canvas!, this.shaderContext, width, + height); } /** Closes and cleans up the resources held by this task. */ diff --git a/mediapipe/tasks/web/vision/image_segmenter/BUILD b/mediapipe/tasks/web/vision/image_segmenter/BUILD index 3db15641..6c1829bd 100644 --- a/mediapipe/tasks/web/vision/image_segmenter/BUILD +++ b/mediapipe/tasks/web/vision/image_segmenter/BUILD @@ -20,7 +20,6 @@ mediapipe_ts_library( "//mediapipe/tasks/cc/vision/image_segmenter/proto:segmenter_options_jspb_proto", "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/vision/core:image_processing_options", - "//mediapipe/tasks/web/vision/core:types", "//mediapipe/tasks/web/vision/core:vision_task_runner", "//mediapipe/util:label_map_jspb_proto", "//mediapipe/web/graph_runner:graph_runner_ts", @@ -36,6 +35,7 @@ mediapipe_ts_declaration( deps = [ "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/core:classifier_options", + "//mediapipe/tasks/web/vision/core:image", "//mediapipe/tasks/web/vision/core:vision_task_options", ], ) @@ -52,6 +52,7 @@ mediapipe_ts_library( "//mediapipe/framework:calculator_jspb_proto", "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/core:task_runner_test_utils", + "//mediapipe/tasks/web/vision/core:image", "//mediapipe/web/graph_runner:graph_runner_image_lib_ts", ], ) diff --git a/mediapipe/tasks/web/vision/image_segmenter/image_segmenter.ts b/mediapipe/tasks/web/vision/image_segmenter/image_segmenter.ts index 4089a2b1..92462815 100644 --- a/mediapipe/tasks/web/vision/image_segmenter/image_segmenter.ts +++ b/mediapipe/tasks/web/vision/image_segmenter/image_segmenter.ts @@ -22,7 +22,6 @@ import {ImageSegmenterGraphOptions as ImageSegmenterGraphOptionsProto} from '../ import {SegmenterOptions as SegmenterOptionsProto} from '../../../../tasks/cc/vision/image_segmenter/proto/segmenter_options_pb'; import {WasmFileset} from '../../../../tasks/web/core/wasm_fileset'; import {ImageProcessingOptions} from '../../../../tasks/web/vision/core/image_processing_options'; -import {SegmentationMask} from '../../../../tasks/web/vision/core/types'; import {VisionGraphRunner, VisionTaskRunner} from '../../../../tasks/web/vision/core/vision_task_runner'; import {LabelMapItem} from '../../../../util/label_map_pb'; import {ImageSource, WasmModule} from '../../../../web/graph_runner/graph_runner'; @@ -33,7 +32,6 @@ import {ImageSegmenterResult} from './image_segmenter_result'; export * from './image_segmenter_options'; export * from './image_segmenter_result'; -export {SegmentationMask}; export {ImageSource}; // Used in the public API const IMAGE_STREAM = 'image_in'; @@ -60,7 +58,7 @@ export type ImageSegmenterCallback = (result: ImageSegmenterResult) => void; /** Performs image segmentation on images. */ export class ImageSegmenter extends VisionTaskRunner { - private result: ImageSegmenterResult = {width: 0, height: 0}; + private result: ImageSegmenterResult = {}; private labels: string[] = []; private outputCategoryMask = DEFAULT_OUTPUT_CATEGORY_MASK; private outputConfidenceMasks = DEFAULT_OUTPUT_CONFIDENCE_MASKS; @@ -313,7 +311,7 @@ export class ImageSegmenter extends VisionTaskRunner { } private reset(): void { - this.result = {width: 0, height: 0}; + this.result = {}; } /** Updates the MediaPipe graph configuration. */ @@ -341,12 +339,8 @@ export class ImageSegmenter extends VisionTaskRunner { this.graphRunner.attachImageVectorListener( CONFIDENCE_MASKS_STREAM, (masks, timestamp) => { - this.result.confidenceMasks = masks.map(m => m.data); - if (masks.length >= 0) { - this.result.width = masks[0].width; - this.result.height = masks[0].height; - } - + this.result.confidenceMasks = + masks.map(wasmImage => this.convertToMPImage(wasmImage)); this.setLatestOutputTimestamp(timestamp); }); this.graphRunner.attachEmptyPacketListener( @@ -361,9 +355,7 @@ export class ImageSegmenter extends VisionTaskRunner { this.graphRunner.attachImageListener( CATEGORY_MASK_STREAM, (mask, timestamp) => { - this.result.categoryMask = mask.data; - this.result.width = mask.width; - this.result.height = mask.height; + this.result.categoryMask = this.convertToMPImage(mask); this.setLatestOutputTimestamp(timestamp); }); this.graphRunner.attachEmptyPacketListener( diff --git a/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_result.d.ts b/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_result.d.ts index b731e032..454ec27e 100644 --- a/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_result.d.ts +++ b/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_result.d.ts @@ -14,24 +14,21 @@ * limitations under the License. */ +import {MPImage} from '../../../../tasks/web/vision/core/image'; + /** The output result of ImageSegmenter. */ export declare interface ImageSegmenterResult { /** - * Multiple masks as Float32Arrays or WebGLTextures where, for each mask, each - * pixel represents the prediction confidence, usually in the [0, 1] range. + * Multiple masks represented as `Float32Array` or `WebGLTexture`-backed + * `MPImage`s where, for each mask, each pixel represents the prediction + * confidence, usually in the [0, 1] range. */ - confidenceMasks?: Float32Array[]|WebGLTexture[]; + confidenceMasks?: MPImage[]; /** - * A category mask as a Uint8ClampedArray or WebGLTexture where each - * pixel represents the class which the pixel in the original image was - * predicted to belong to. + * A category mask represented as a `Uint8ClampedArray` or + * `WebGLTexture`-backed `MPImage` where each pixel represents the class which + * the pixel in the original image was predicted to belong to. */ - categoryMask?: Uint8ClampedArray|WebGLTexture; - - /** The width of the masks. */ - width: number; - - /** The height of the masks. */ - height: number; + categoryMask?: MPImage; } diff --git a/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_test.ts b/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_test.ts index 1327c1b3..7f7f1906 100644 --- a/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_test.ts +++ b/mediapipe/tasks/web/vision/image_segmenter/image_segmenter_test.ts @@ -20,6 +20,7 @@ import 'jasmine'; import {CalculatorGraphConfig} from '../../../../framework/calculator_pb'; import {addJasmineCustomFloatEqualityTester, createSpyWasmModule, MediapipeTasksFake, SpyWasmModule, verifyGraph} from '../../../../tasks/web/core/task_runner_test_utils'; import {WasmImage} from '../../../../web/graph_runner/graph_runner_image_lib'; +import {MPImage} from '../../../../tasks/web/vision/core/image'; import {ImageSegmenter} from './image_segmenter'; import {ImageSegmenterOptions} from './image_segmenter_options'; @@ -182,10 +183,10 @@ describe('ImageSegmenter', () => { return new Promise(resolve => { imageSegmenter.segment({} as HTMLImageElement, result => { expect(imageSegmenter.fakeWasmModule._waitUntilIdle).toHaveBeenCalled(); - expect(result.categoryMask).toEqual(mask); + expect(result.categoryMask).toBeInstanceOf(MPImage); expect(result.confidenceMasks).not.toBeDefined(); - expect(result.width).toEqual(2); - expect(result.height).toEqual(2); + expect(result.categoryMask!.width).toEqual(2); + expect(result.categoryMask!.height).toEqual(2); resolve(); }); }); @@ -214,18 +215,21 @@ describe('ImageSegmenter', () => { imageSegmenter.segment({} as HTMLImageElement, result => { expect(imageSegmenter.fakeWasmModule._waitUntilIdle).toHaveBeenCalled(); expect(result.categoryMask).not.toBeDefined(); - expect(result.confidenceMasks).toEqual([mask1, mask2]); - expect(result.width).toEqual(2); - expect(result.height).toEqual(2); + + expect(result.confidenceMasks![0]).toBeInstanceOf(MPImage); + expect(result.confidenceMasks![0].width).toEqual(2); + expect(result.confidenceMasks![0].height).toEqual(2); + + expect(result.confidenceMasks![1]).toBeInstanceOf(MPImage); resolve(); }); }); }); it('supports combined category and confidence masks', async () => { - const categoryMask = new Uint8ClampedArray([1, 0]); - const confidenceMask1 = new Float32Array([0.0, 1.0]); - const confidenceMask2 = new Float32Array([1.0, 0.0]); + const categoryMask = new Uint8ClampedArray([1]); + const confidenceMask1 = new Float32Array([0.0]); + const confidenceMask2 = new Float32Array([1.0]); await imageSegmenter.setOptions( {outputCategoryMask: true, outputConfidenceMasks: true}); @@ -248,12 +252,12 @@ describe('ImageSegmenter', () => { // Invoke the image segmenter imageSegmenter.segment({} as HTMLImageElement, result => { expect(imageSegmenter.fakeWasmModule._waitUntilIdle).toHaveBeenCalled(); - expect(result.categoryMask).toEqual(categoryMask); - expect(result.confidenceMasks).toEqual([ - confidenceMask1, confidenceMask2 - ]); - expect(result.width).toEqual(1); - expect(result.height).toEqual(1); + expect(result.categoryMask).toBeInstanceOf(MPImage); + expect(result.categoryMask!.width).toEqual(1); + expect(result.categoryMask!.height).toEqual(1); + + expect(result.confidenceMasks![0]).toBeInstanceOf(MPImage); + expect(result.confidenceMasks![1]).toBeInstanceOf(MPImage); resolve(); }); });