diff --git a/js/web/lib/backend-wasm.ts b/js/web/lib/backend-wasm.ts index 21a759557e..961b331d77 100644 --- a/js/web/lib/backend-wasm.ts +++ b/js/web/lib/backend-wasm.ts @@ -1,10 +1,8 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -import {readFile} from 'fs'; import {Backend, env, InferenceSession, SessionHandler} from 'onnxruntime-common'; import {cpus} from 'os'; -import {promisify} from 'util'; import {initWasm} from './wasm/proxy-wrapper'; import {OnnxruntimeWebAssemblySessionHandler} from './wasm/session-handler'; @@ -46,23 +44,8 @@ class OnnxruntimeWebAssemblyBackend implements Backend { createSessionHandler(buffer: Uint8Array, options?: InferenceSession.SessionOptions): Promise; async createSessionHandler(pathOrBuffer: string|Uint8Array, options?: InferenceSession.SessionOptions): Promise { - let buffer: Uint8Array; - if (typeof pathOrBuffer === 'string') { - if (typeof fetch === 'undefined') { - // node - buffer = await promisify(readFile)(pathOrBuffer); - } else { - // browser - const response = await fetch(pathOrBuffer); - const arrayBuffer = await response.arrayBuffer(); - buffer = new Uint8Array(arrayBuffer); - } - } else { - buffer = pathOrBuffer; - } - const handler = new OnnxruntimeWebAssemblySessionHandler(); - await handler.loadModel(buffer, options); + await handler.loadModel(pathOrBuffer, options); return Promise.resolve(handler); } } diff --git a/js/web/lib/wasm/proxy-messages.ts b/js/web/lib/wasm/proxy-messages.ts index d8dcbd3664..7abd7062e0 100644 --- a/js/web/lib/wasm/proxy-messages.ts +++ b/js/web/lib/wasm/proxy-messages.ts @@ -13,6 +13,11 @@ export type SerializableTensor = [Tensor.Type, readonly number[], Tensor.DataTyp */ export type SerializableSessionMetadata = [number, string[], string[]]; +/** + * tuple elements are: modeldata.offset, modeldata.length + */ +export type SerializableModeldata = [number, number]; + interface MessageError { err?: string; } @@ -27,6 +32,18 @@ interface MessageInitOrt extends MessageError { in ?: {numThreads: number; loggingLevel: number}; } +interface MessageCreateSessionAllocate extends MessageError { + type: 'create_allocate'; + in ?: {model: Uint8Array}; + out?: SerializableModeldata; +} + +interface MessageCreateSessionFinalize extends MessageError { + type: 'create_finalize'; + in ?: {modeldata: SerializableModeldata; options?: InferenceSession.SessionOptions}; + out?: SerializableSessionMetadata; +} + interface MessageCreateSession extends MessageError { type: 'create'; in ?: {model: Uint8Array; options?: InferenceSession.SessionOptions}; @@ -52,5 +69,5 @@ interface MesssageEndProfiling extends MessageError { in ?: number; } -export type OrtWasmMessage = - MessageInitWasm|MessageInitOrt|MessageCreateSession|MessageReleaseSession|MessageRun|MesssageEndProfiling; +export type OrtWasmMessage = MessageInitWasm|MessageInitOrt|MessageCreateSessionAllocate|MessageCreateSessionFinalize| + MessageCreateSession|MessageReleaseSession|MessageRun|MesssageEndProfiling; diff --git a/js/web/lib/wasm/proxy-worker/main.ts b/js/web/lib/wasm/proxy-worker/main.ts index 1e343f890a..b72bfe42c6 100644 --- a/js/web/lib/wasm/proxy-worker/main.ts +++ b/js/web/lib/wasm/proxy-worker/main.ts @@ -4,7 +4,7 @@ /// import {OrtWasmMessage} from '../proxy-messages'; -import {createSession, endProfiling, extractTransferableBuffers, initOrt, releaseSession, run} from '../wasm-core-impl'; +import {createSession, createSessionAllocate, createSessionFinalize, endProfiling, extractTransferableBuffers, initOrt, releaseSession, run} from '../wasm-core-impl'; import {initializeWebAssembly} from '../wasm-factory'; self.onmessage = (ev: MessageEvent): void => { @@ -24,6 +24,24 @@ self.onmessage = (ev: MessageEvent): void => { postMessage({type: 'init-ort', err} as OrtWasmMessage); } break; + case 'create_allocate': + try { + const {model} = ev.data.in!; + const modeldata = createSessionAllocate(model); + postMessage({type: 'create_allocate', out: modeldata} as OrtWasmMessage); + } catch (err) { + postMessage({type: 'create_allocate', err} as OrtWasmMessage); + } + break; + case 'create_finalize': + try { + const {modeldata, options} = ev.data.in!; + const sessionMetadata = createSessionFinalize(modeldata, options); + postMessage({type: 'create_finalize', out: sessionMetadata} as OrtWasmMessage); + } catch (err) { + postMessage({type: 'create_finalize', err} as OrtWasmMessage); + } + break; case 'create': try { const {model, options} = ev.data.in!; diff --git a/js/web/lib/wasm/proxy-wrapper.ts b/js/web/lib/wasm/proxy-wrapper.ts index d50cc1449f..75664ac4af 100644 --- a/js/web/lib/wasm/proxy-wrapper.ts +++ b/js/web/lib/wasm/proxy-wrapper.ts @@ -3,7 +3,7 @@ import {env, InferenceSession} from 'onnxruntime-common'; -import {OrtWasmMessage, SerializableSessionMetadata, SerializableTensor} from './proxy-messages'; +import {OrtWasmMessage, SerializableModeldata, SerializableSessionMetadata, SerializableTensor} from './proxy-messages'; import * as core from './wasm-core-impl'; import {initializeWebAssembly} from './wasm-factory'; @@ -18,6 +18,8 @@ type PromiseCallbacks = [(result: T) => void, (reason: unknown) => voi let initWasmCallbacks: PromiseCallbacks; let initOrtCallbacks: PromiseCallbacks; +const createSessionAllocateCallbacks: Array> = []; +const createSessionFinalizeCallbacks: Array> = []; const createSessionCallbacks: Array> = []; const releaseSessionCallbacks: Array> = []; const runCallbacks: Array> = []; @@ -48,6 +50,20 @@ const onProxyWorkerMessage = (ev: MessageEvent): void => { initOrtCallbacks[0](); } break; + case 'create_allocate': + if (ev.data.err) { + createSessionAllocateCallbacks.shift()![1](ev.data.err); + } else { + createSessionAllocateCallbacks.shift()![0](ev.data.out!); + } + break; + case 'create_finalize': + if (ev.data.err) { + createSessionFinalizeCallbacks.shift()![1](ev.data.err); + } else { + createSessionFinalizeCallbacks.shift()![0](ev.data.out!); + } + break; case 'create': if (ev.data.err) { createSessionCallbacks.shift()![1](ev.data.err); @@ -131,6 +147,33 @@ export const initOrt = async(numThreads: number, loggingLevel: number): Promise< } }; +export const createSessionAllocate = async(model: Uint8Array): Promise => { + if (!BUILD_DEFS.DISABLE_WASM_PROXY && isProxy()) { + ensureWorker(); + return new Promise((resolve, reject) => { + createSessionAllocateCallbacks.push([resolve, reject]); + const message: OrtWasmMessage = {type: 'create_allocate', in : {model}}; + proxyWorker!.postMessage(message, [model.buffer]); + }); + } else { + return core.createSessionAllocate(model); + } +}; + +export const createSessionFinalize = async(modeldata: SerializableModeldata, options?: InferenceSession.SessionOptions): + Promise => { + if (!BUILD_DEFS.DISABLE_WASM_PROXY && isProxy()) { + ensureWorker(); + return new Promise((resolve, reject) => { + createSessionFinalizeCallbacks.push([resolve, reject]); + const message: OrtWasmMessage = {type: 'create_finalize', in : {modeldata, options}}; + proxyWorker!.postMessage(message); + }); + } else { + return core.createSessionFinalize(modeldata, options); + } + }; + export const createSession = async(model: Uint8Array, options?: InferenceSession.SessionOptions): Promise => { if (!BUILD_DEFS.DISABLE_WASM_PROXY && isProxy()) { diff --git a/js/web/lib/wasm/session-handler.ts b/js/web/lib/wasm/session-handler.ts index ede51e8a24..a507b09e89 100644 --- a/js/web/lib/wasm/session-handler.ts +++ b/js/web/lib/wasm/session-handler.ts @@ -1,9 +1,12 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +import {readFile} from 'fs'; import {env, InferenceSession, SessionHandler, Tensor} from 'onnxruntime-common'; +import {promisify} from 'util'; -import {createSession, endProfiling, initOrt, releaseSession, run} from './proxy-wrapper'; +import {SerializableModeldata} from './proxy-messages'; +import {createSession, createSessionAllocate, createSessionFinalize, endProfiling, initOrt, releaseSession, run} from './proxy-wrapper'; let ortInit: boolean; @@ -25,19 +28,42 @@ const getLogLevel = (logLevel: 'verbose'|'info'|'warning'|'error'|'fatal'): numb } }; + export class OnnxruntimeWebAssemblySessionHandler implements SessionHandler { private sessionId: number; inputNames: string[]; outputNames: string[]; - async loadModel(model: Uint8Array, options?: InferenceSession.SessionOptions): Promise { + async createSessionAllocate(path: string): Promise { + // fetch model from url and move to wasm heap. The arraybufffer that held the http + // response is freed once we return + const response = await fetch(path); + const arrayBuffer = await response.arrayBuffer(); + return createSessionAllocate(new Uint8Array(arrayBuffer)); + } + + async loadModel(pathOrBuffer: string|Uint8Array, options?: InferenceSession.SessionOptions): Promise { if (!ortInit) { await initOrt(env.wasm.numThreads!, getLogLevel(env.logLevel!)); ortInit = true; } - [this.sessionId, this.inputNames, this.outputNames] = await createSession(model, options); + if (typeof pathOrBuffer === 'string') { + if (typeof fetch === 'undefined') { + // node + const model = await promisify(readFile)(pathOrBuffer); + [this.sessionId, this.inputNames, this.outputNames] = await createSession(model, options); + } else { + // browser + // fetch model and move to wasm heap. + const modelData: SerializableModeldata = await this.createSessionAllocate(pathOrBuffer); + // create the session + [this.sessionId, this.inputNames, this.outputNames] = await createSessionFinalize(modelData, options); + } + } else { + [this.sessionId, this.inputNames, this.outputNames] = await createSession(pathOrBuffer, options); + } } async dispose(): Promise { diff --git a/js/web/lib/wasm/wasm-core-impl.ts b/js/web/lib/wasm/wasm-core-impl.ts index 890ef58ada..4406f09714 100644 --- a/js/web/lib/wasm/wasm-core-impl.ts +++ b/js/web/lib/wasm/wasm-core-impl.ts @@ -3,7 +3,7 @@ import {InferenceSession, Tensor} from 'onnxruntime-common'; -import {SerializableSessionMetadata, SerializableTensor} from './proxy-messages'; +import {SerializableModeldata, SerializableSessionMetadata, SerializableTensor} from './proxy-messages'; import {setRunOptions} from './run-options'; import {setSessionOptions} from './session-options'; import {allocWasmString} from './string-utils'; @@ -32,10 +32,17 @@ const activeSessions = new Map(); * create an instance of InferenceSession. * @returns the metadata of InferenceSession. 0-value handle for failure. */ -export const createSession = - (model: Uint8Array, options?: InferenceSession.SessionOptions): SerializableSessionMetadata => { +export const createSessionAllocate = (model: Uint8Array): [number, number] => { + const wasm = getInstance(); + const modelDataOffset = wasm._malloc(model.byteLength); + wasm.HEAPU8.set(model, modelDataOffset); + return [modelDataOffset, model.byteLength]; +}; + +export const createSessionFinalize = + (modelData: SerializableModeldata, options?: InferenceSession.SessionOptions): SerializableSessionMetadata => { const wasm = getInstance(); - const modelDataOffset = wasm._malloc(model.byteLength); + let sessionHandle = 0; let sessionOptionsHandle = 0; let allocs: number[] = []; @@ -43,13 +50,12 @@ export const createSession = try { [sessionOptionsHandle, allocs] = setSessionOptions(options); - wasm.HEAPU8.set(model, modelDataOffset); - sessionHandle = wasm._OrtCreateSession(modelDataOffset, model.byteLength, sessionOptionsHandle); + sessionHandle = wasm._OrtCreateSession(modelData[0], modelData[1], sessionOptionsHandle); if (sessionHandle === 0) { throw new Error('Can\'t create a session'); } } finally { - wasm._free(modelDataOffset); + wasm._free(modelData[0]); wasm._OrtReleaseSessionOptions(sessionOptionsHandle); allocs.forEach(wasm._free); } @@ -82,6 +88,17 @@ export const createSession = return [sessionHandle, inputNames, outputNames]; }; + +/** + * create an instance of InferenceSession. + * @returns the metadata of InferenceSession. 0-value handle for failure. + */ +export const createSession = + (model: Uint8Array, options?: InferenceSession.SessionOptions): SerializableSessionMetadata => { + const modelData: SerializableModeldata = createSessionAllocate(model); + return createSessionFinalize(modelData, options); + }; + export const releaseSession = (sessionId: number): void => { const wasm = getInstance(); const session = activeSessions.get(sessionId);