Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -9,14 +9,20 @@ import { useModuleFactory } from '../useModuleFactory';
* @returns Ready to use VAD model.
*/
export const useVAD = ({ model, preventLoad = false }: VADProps): VADType => {
const { error, isReady, isGenerating, downloadProgress, runForward } =
useModuleFactory({
factory: (config, onProgress) =>
VADModule.fromModelName(config, onProgress),
config: model,
deps: [model.modelName, model.modelSource],
preventLoad,
});
const {
error,
isReady,
isGenerating,
downloadProgress,
runForward,
runSideChannel,
} = useModuleFactory({
factory: (config, onProgress) =>
VADModule.fromModelName(config, onProgress),
config: model,
deps: [model.modelName, model.modelSource],
preventLoad,
});

const forward = (waveform: Float32Array) =>
runForward((inst) => inst.forward(waveform));
Expand All @@ -25,16 +31,9 @@ export const useVAD = ({ model, preventLoad = false }: VADProps): VADType => {
runForward((inst) => inst.stream(input));

const streamInsert = (waveform: Float32Array) =>
runForward((inst) => {
inst.streamInsert(waveform);
return Promise.resolve();
});
runSideChannel((inst) => inst.streamInsert(waveform));

const streamStop = () =>
runForward((inst) => {
inst.streamStop();
return Promise.resolve();
});
const streamStop = () => runSideChannel((inst) => inst.streamStop());

return {
error,
Expand Down
13 changes: 12 additions & 1 deletion packages/react-native-executorch/src/hooks/useModuleFactory.ts
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ type RunOnFrame<M> = M extends { runOnFrame: infer R } ? R : never;
* not-loaded / already-generating guards so individual hooks only need to
* define their typed `forward` wrapper.
* @param props - Options object containing the factory function, config, deps array, and optional preventLoad flag.
* @returns An object with error, isReady, isGenerating, downloadProgress, runForward, instance, and runOnFrame.
* @returns An object with error, isReady, isGenerating, downloadProgress, runForward, runSideChannel, instance, and runOnFrame.
* @internal
*/
export function useModuleFactory<M extends Deletable, Config>({
Expand Down Expand Up @@ -89,6 +89,16 @@ export function useModuleFactory<M extends Deletable, Config>({
}
};

// Non-gating call path for streaming modules: only checks `isReady`, never
// `isGenerating`. Lets side-channel methods (e.g. `streamInsert` buffer push,
// `streamStop` interrupt signal) run while `stream`/`forward` is in flight.
const runSideChannel = <R>(fn: (instance: M) => R): R => {
if (!isReady || !instance) {
throw new RnExecutorchError(RnExecutorchErrorCode.ModuleNotLoaded);
}
return fn(instance);
};

const runOnFrame = useMemo(
() =>
instance && 'runOnFrame' in instance
Expand All @@ -103,6 +113,7 @@ export function useModuleFactory<M extends Deletable, Config>({
isGenerating,
downloadProgress,
runForward,
runSideChannel,
instance,
runOnFrame,
};
Expand Down
Loading