Trois champs de mines rencontrés lors de la migration de TensorFlow.js vers LiteRT.js
Faire tourner des modèles d'IA dans un navigateur web est passionnant, mais on atteint vite ses limites. Dès que l'on traite des images haute résolution ou que l'on lance des calculs lourds, l'affichage saccade. TensorFlow.js (TF.js) utilise le backend WebGL et des liaisons de noyau JavaScript, ce qui crée une surcharge importante lors des calculs matriciels à grande échelle.
En revanche, LiteRT.js compile un runtime natif C++ en WebAssembly (Wasm) pour l'intégrer au navigateur. C'est une structure qui résout les goulots d'étranglement à la racine. De plus, le processus d'importation de modèles PyTorch vers le web est simplifié. Auparavant, il fallait passer de PyTorch à ONNX, puis via TensorFlow jusqu'au format TF.js, entraînant des ruptures de compatibilité des opérateurs et des pertes de précision. Désormais, une seule bibliothèque (ai-edge-torch) suffit pour extraire un fichier standard .tflite.
Si vous ne voulez pas abandonner vos actifs existants, vous pouvez utiliser le paquet @litertjs/tfjs-interop fourni par Google. La méthode consiste à conserver le pipeline de prétraitement ou de post-traitement de données de TF.js tel quel, et à remplacer uniquement la partie exécution de prédiction du modèle par LiteRT.js.
Résoudre l'incompatibilité entre NCHW et NHWC
Les modèles TFLite convertis depuis PyTorch exigent généralement une structure canal prioritaire NCHW (Canal, Hauteur, Largeur). Cependant, le tableau ImageData du Canvas du navigateur est au format NHWC (Hauteur, Largeur, Canal) avec les pixels alignés. Il est nécessaire de disposer d'un code de prétraitement qui traite cet écart de manière synchrone sans surcharger le thread principal.
Tout d'abord, on multiplie le nombre total de pixels de l'ImageData par 3 pour créer un Float32Array. Ensuite, on spécifie les positions des offsets de mise à plat pour chaque canal rouge, vert et bleu à 0, au nombre total de pixels, et au double du nombre total de pixels. Enfin, on normalise les valeurs des pixels de 0 à 255 en divisant par 255.0, puis on les assigne à chaque offset de canal. C'est le processus de réorganisation des données d'entrée en un tampon de tenseur aplati NCHW.
`javascript
/**
- Utilitaire de prétraitement pour une conversion rapide d'un tampon ImageData au format NHWC vers NCHW Float32Array haute performance
- @param {ImageData} imageData - Données de pixels brutes obtenues depuis un Canvas HTML5
- @param {number} width - Résolution horizontale de l'image d'entrée requise par le modèle cible
- @param {number} height - Résolution verticale de l'image d'entrée requise par le modèle cible
- @returns {Float32Array} Tampon de tenseur aplati réorganisé au format NCHW
*/
export function preprocessNHWCToNCHW(imageData, width, height) {
const { data } = imageData;
const totalPixels = width * height;
const nchwBuffer = new Float32Array(totalPixels * 3);
const rChannelOffset = 0;
const gChannelOffset = totalPixels;
const bChannelOffset = totalPixels * 2;
for (let i = 0; i < totalPixels; i++) {
const srcIndex = i * 4;
nchwBuffer[rChannelOffset + i] = data[srcIndex] / 255.0;
nchwBuffer[gChannelOffset + i] = data[srcIndex + 1] / 255.0;
nchwBuffer[bChannelOffset + i] = data[srcIndex + 2] / 255.0;
}
return nchwBuffer;
}
`
Éviter les saccades d'écran avec les Web Workers et le Zero-Copy
La situation la plus critique lors de l'exécution de deep learning sur le frontend est le blocage de l'interface utilisateur. Pour qu'un navigateur puisse afficher une animation fluide de 60 images par seconde, la boucle d'événements doit traiter les tâches synchrones en moins de 16,6 ms. Or, les calculs de tenseurs bloquent facilement la boucle de rendu principale, qui est mono-thread.
LiteRT.js a une vitesse d'exécution de base environ 3 fois plus rapide que les outils basés sur JavaScript. Avec l'accélération WebGPU ou WebNN, il devient de 5 à 60 fois plus rapide que le mode CPU. Le backend WebNN, qui utilise une NPU dédiée, exige impérativement l'activation de l'intégration des promesses JavaScript (JSPI) pour relier le planificateur de noyau WebAssembly synchrone à la boucle de contrôle matérielle asynchrone du navigateur. Pour utiliser ces ressources d'accélération tout en préservant le thread principal, il est indispensable d'isoler l'initialisation de la bibliothèque et l'intégralité du pipeline d'inférence au sein d'un Web Worker.
À ce stade, si l'on transmet simplement les tampons lors de la communication de données entre threads, une copie mémoire interne se produit, accumulant une surcharge sur le CPU et la mémoire heap. Il faut utiliser des objets transférables (Transferable Objects) pour transférer la propriété même de la zone d'adresse mémoire physique afin d'éliminer la latence. Les tampons dont la propriété a été transférée sont immédiatement invalidés dans le contexte source, garantissant ainsi la sécurité entre les threads.
`javascript
// litert-worker.js - Module Web Worker dédié aux calculs d'inférence en arrière-plan
import { loadLiteRt, loadAndCompile, Tensor } from '@litertjs/core';
let compiledModel = null;
let isLoaded = false;
self.onmessage = async (event) => {
const { type, payload } = event.data;
switch (type) {
case 'LOAD_MODEL':
try {
await loadLiteRt(payload.wasmDirectory, { jspi: payload.enableJspi || false });
compiledModel = await loadAndCompile(payload.modelUrl, {
accelerator: payload.accelerator || 'webgpu'
});
isLoaded = true;
self.postMessage({ type: 'MODEL_READY' });
} catch (err) {
self.postMessage({ type: 'ERROR', error: Initialization failed: ${err.message} });
}
break;
case 'RUN_INFERENCE':
if (!isLoaded || !compiledModel) {
self.postMessage({ type: 'ERROR', error: 'Model has not been loaded' });
return;
}
try {
const rawInputData = payload.bufferData;
const inputShape = payload.shape;
const inputTensor = new Tensor(rawInputData, inputShape);
const results = await compiledModel.run(inputTensor);
const cpuOutputTensor = await results[0].moveTo('wasm');
const outputBuffer = cpuOutputTensor.toTypedArray();
inputTensor.delete();
cpuOutputTensor.delete();
results[0].delete();
self.postMessage(
{
type: 'INFERENCE_COMPLETE',
payload: {
data: outputBuffer,
shape: results[0].shape
}
},
[outputBuffer.buffer]
);
} catch (err) {
self.postMessage({ type: 'ERROR', error: `Inference failed: ${err.message}` });
}
break;
default:
self.postMessage({ type: 'UNKNOWN_OP' });
}
};
`
`javascript
// litert-bridge.js - Classe orchestratrice d'IA pour le thread principal
export class LiteRtBridge {
constructor(workerPath) {
this.worker = new Worker(workerPath);
this.promiseMap = new Map();
this.tokenCounter = 0;
this.worker.onmessage = (event) => {
const { type, payload, error } = event.data;
if (type === 'MODEL_READY') {
if (this.initResolve) this.initResolve();
} else if (type === 'INFERENCE_COMPLETE') {
const currentToken = this.tokenCounter;
const promiseHandler = this.promiseMap.get(currentToken);
if (promiseHandler) {
promiseHandler.resolve(payload);
this.promiseMap.delete(currentToken);
}
} else if (type === 'ERROR') {
const currentToken = this.tokenCounter;
const promiseHandler = this.promiseMap.get(currentToken);
if (promiseHandler) {
promiseHandler.reject(new Error(error));
this.promiseMap.delete(currentToken);
} else if (this.initReject) {
this.initReject(new Error(error));
}
}
};
}
bootstrap(wasmDirectory, modelUrl, accelerator = 'webgpu') {
return new Promise((resolve, reject) => {
this.initResolve = resolve;
this.initReject = reject;
this.worker.postMessage({
type: 'LOAD_MODEL',
payload: { wasmDirectory, modelUrl, accelerator, enableJspi: true }
});
});
}
execute(inputFloat32Array, inputShape) {
return new Promise((resolve, reject) => {
this.tokenCounter++;
this.promiseMap.set(this.tokenCounter, { resolve, reject });
this.worker.postMessage(
{
type: 'RUN_INFERENCE',
payload: {
bufferData: inputFloat32Array,
shape: inputShape
}
},
[inputFloat32Array.buffer]
);
});
}
}
`
Ne faites pas confiance au ramasse-miettes automatique
Lors de l'utilisation de TF.js, le modèle standard consistait à supprimer les tenseurs dans une portée d'appel synchrone avec tf.tidy(). Cependant, si du code asynchrone ou des promesses étaient impliqués, cela créait des bugs où les tenseurs étaient supprimés alors que le travail asynchrone n'était pas terminé, ou bien n'étaient pas collectés du tout.
LiteRT.js est encore plus impitoyable. Il n'est pas soumis au ramasse-miettes (GC) du moteur du navigateur. Le moteur JavaScript, comme V8, ne peut pas suivre l'état de la heap dans l'espace mémoire virtuel linéaire de WebAssembly ni dans les tampons WebGPU. Si vous n'appelez pas explicitement .delete() sur une instance de tenseur après usage, la mémoire du navigateur s'accumulera à l'infini. Pour un service diffusant des images vidéo haute définition des dizaines de fois par seconde, l'onglet finira par planter en quelques minutes.
Il est préférable de créer une classe de suivi de portée (scope tracker) qui enregistre la durée de vie des tenseurs créés dans l'ensemble du pipeline asynchrone et garantit leur destruction collective, afin de gérer cela manuellement en toute tranquillité.
`javascript
/**
- Gestionnaire de portée mémoire asynchrone qui assure le suivi manuel et la destruction certaine des tenseurs heap WebAssembly
*/
export class LiteRtScopeTracker {
constructor() {
this.trackList = new Set();
}
/**
- Intégrer les tenseurs créés ou transférés dans la liste de gestion du cycle de vie
- @param {Tensor} tensor - Tenseur LiteRT.js dont on veut assurer le suivi et la destruction
- @returns {Tensor} Retourne l'objet tenseur reçu pour permettre une écriture de code en ligne
*/
register(tensor) {
if (tensor && typeof tensor.delete === 'function') {
this.trackList.add(tensor);
}
return tensor;
}
/**
- Forcer une structure de gestion de pipeline de tenseur sécurisée au sein d'un bloc d'exécution asynchrone
- @param {Function} asyncCallable - Fonction de logique métier d'inférence asynchrone
- @returns {Promise<*>} Résultat final de données brutes retourné par le bloc d'exécution
*/
async enforceScope(asyncCallable) {
try {
const outputResult = await asyncCallable(this);
if (Array.isArray(outputResult)) {
outputResult.forEach((item) => this.trackList.delete(item));
} else {
this.trackList.delete(outputResult);
}
return outputResult;
} finally {
this.disposeAll();
}
}
/**
- Libère définitivement de la zone Wasm tous les tenseurs TFLite natifs liés et encore actifs
*/
disposeAll() {
for (const tensor of this.trackList) {
try {
tensor.delete();
} catch (err) {
console.error('An error occurred while cleaning the native Wasm tensor memory:', err);
}
}
this.trackList.clear();
}
}
`
`javascript
// Exemple de mise en œuvre d'un traitement d'inférence IA asynchrone multiple, sûr et robuste, utilisant le scope tracker mémoire
export async function runRobustVisionInference(rawPixelArray, compiledModel) {
const scopeTracker = new LiteRtScopeTracker();
try {
return await scopeTracker.enforceScope(async (scope) => {
const inputTensor = scope.register(new Tensor(rawPixelArray, [1, 3, 224, 224]));
const predictionResults = await compiledModel.run(inputTensor);
predictionResults.forEach((tensor) => scope.register(tensor));
const firstOutputTensor = predictionResults[0];
const wasmTransferTensor = scope.register(await firstOutputTensor.moveTo('wasm'));
const targetJsArray = wasmTransferTensor.toTypedArray();
return targetJsArray;
});
} catch (err) {
console.error('Fatal crash occurred during the model pipeline execution:', err);
throw err;
}
}
`
Charger les fichiers Wasm lourds de manière conditionnelle
Le problème de l'augmentation de la taille du téléchargement des ressources ne peut pas non plus être ignoré. Pour que le tree-shaking fonctionne dans le bundler, il faut supprimer les références statiques de style CommonJS et rédiger le code source sur la base de la syntaxe des modules ES6 (import/export). Il faut également veiller à configurer sideEffects: false dans l'outil de build pour alléger le bundle.
Le runtime de base de LiteRT.js, @litertjs/core, charge sélectivement trois versions de noyaux WebAssembly en fonction des performances de l'appareil. Sur des navigateurs récents comme Chrome ou Edge, il sélectionne un module prenant en charge le multi-threading et SIMD (litert_wasm_simd.wasm), et pour les environnements hérités comme Safari, il prend le module de repli par défaut (litert_wasm.wasm). Si la compilation GPU échoue, le runtime XNNPACK, qui déporte tous les opérateurs matériels vers la sandbox Wasm du CPU, agit comme dispositif auxiliaire.
Pour éviter les retards de chargement initial, il faut vérifier l'accélérateur spécifique à la configuration et importer dynamiquement le module.
`javascript
// litert-loader.js - Moteur de découverte de périphérique runtime et de couplage dynamique d'accélérateur
export async function bootstrapHighPerformanceInferenceEngine() {
const supportsWebGpu = 'gpu' in navigator;
let chosenAccelerator = 'wasm';
if (supportsWebGpu) {
try {
const gpuAdapter = await navigator.gpu.requestAdapter();
if (gpuAdapter) {
const info = await gpuAdapter.requestDevice();
if (info) {
chosenAccelerator = 'webgpu';
}
}
} catch (e) {
console.warn("GPU profile probe failed, resolving execution chain to fallback WASM.");
}
}
const { loadLiteRt, loadAndCompile } = await import('@litertjs/core');
const cdnWasmHostPath = 'https://cdn.jsdelivr.net/npm/@litertjs/core/wasm/';
await loadLiteRt(cdnWasmHostPath, {
jspi: chosenAccelerator === 'webnn'
});
return {
loadAndCompile,
chosenAccelerator
};
}
`
Intégrez le transfert de propriété de la mémoire via les Web Workers et la destruction explicite des objets dans la zone Wasm dans vos principes de conception. Une fois que vous commencez à contrôler directement les flux de données, vous pourrez déployer des services d'IA embarqués en production sans craindre que le navigateur n'explose.