tf.js tfjs-automl 实践指南:AutoML Edge 图像分类与目标检测推理 API tf.js tfjs-automl 实践指南AutoML Edge 图像分类与目标检测推理 API【免费下载链接】tfjsA WebGL accelerated JavaScript library for training and deploying ML models.项目地址: https://gitcode.com/gh_mirrors/tf/tfjstfjs-automl 是 TensorFlow.js 仓库中面向 AutoML Edge 产物的推理封装包它把加载模型文件 → 图像预处理 → 模型推理 → 结果后处理整条链路收敛为几个可直接调用的 API。读完本篇你将掌握如何在 Web 或非浏览器环境中通过loadImageClassification/loadObjectDetection加载 AutoML Edge 训练出的图像分类与目标检测模型理解classify/detect的参数语义centerCrop、score、iou、topk及其在 src/img_classification.ts、src/object_detection.ts 中的实现细节并能直接运行仓库内的两个官方 Demo。包定位与安装tfjs-automl 提供的是一组加载并运行 AutoML Edge 模型 API包名tensorflow/tfjs-automl见 tfjs-automl/package.json。它支持两类模型任务图像分类Image classification目标检测Object detection安装方式与 tf.js 家族其他包一致npm / yarnnpm i tensorflow/tfjs-automlCDN引入后以全局方式使用script srchttps://cdn.jsdelivr.net/npm/tensorflow/tfjs-automl/script从 tfjs-automl/package.json 的peerDependencies可以看出它依赖tensorflow/tfjs-core、tensorflow/tfjs-converter与tensorflow/tfjs-backend-webgl均要求 ^3.9.0 以上这与官方文档的提示一致推理需要核心运行时张量计算、converterGraphModel 加载以及一个后端浏览器下通常为 WebGL。公共 API 从 src/index.ts 统一导出// Image classification API. export {ImageClassificationModel, ImageClassificationOptions, ImagePrediction, loadImageClassification} from ./img_classification; // Object detection API. export {Box, loadObjectDetection, ObjectDetectionModel, ObjectDetectionOptions, PredictedObject} from ./object_detection; // Shared API. export {ImageInput} from ./types; export {version} from ./version;AutoML Edge 模型的三件套文件两类模型的产物文件结构相同这是使用 AutoML Edge 模型的前置条件model.json模型拓扑graph modeldict.txt以换行符分隔的标签列表newline-separated labels一个或多个*.bin文件承载权重。文档强调必须确保这些文件可以作为静态资源被 Web 应用访问——本地起服务或托管到 Google Cloud Storage 均可。这一点从源码实现上得到印证loadImageClassification(modelUrl)内部对modelUrl指向model.json的 URL与dict.txt分别发起并行请求dict.txt被解析为有序字符串数组后与GraphModel一起注入模型实例见 src/img_classification.ts 中的loadImageClassification与 src/util.ts 中的loadDictionary。loadDictionary的 URL 拼接逻辑值得注意它取modelUrl中最后一个/之前的前缀在其后拼接dict.txt。因此无论传入model.json、base/path/model.json还是绝对路径/model.json字典都会被解析到与模型同目录的位置。这一行为在 src/util_test.ts 中有针对性测试四种相对/绝对路径组合分别断言请求 URL 为dict.txt、base/path/dict.txt、/dict.txt、/base/path/dict.txt并验证按换行切分、trim()后返回[first, second, third]。图像分类运行官方 Demo图像分类 Demo 位于 demo/img_classification运行步骤cd demo/img_classification yarn yarn watchyarn watch会启动一个本地 HTTP 服务器端口1234并托管 Demo 页面。从该目录的 package.json 可以看到watch脚本实际调用 Parcelcross-env NODE_ENVdevelopment parcel index.html --no-hmr --open且依赖中通过tensorflow/tfjs-automl: link:../../以本地链接方式引用本仓库的包便于联调。加载模型HTTP 方式加载最常用import * as automl from tensorflow/tfjs-automl; const modelUrl model.json; // URL to the model.json file. const model await automl.loadImageClassification(modelUrl);如果不想或不能通过 HTTP 加载模型——这在非浏览器平台尤其常见——可以分别加载 graph model 与字典然后直接使用构造函数import * as automl from tensorflow/tfjs-automl; import * as tf from tensorflow/tfjs; // You can load the graph model using any IO handler const graphModel await tf.loadGraphModel(string|io.IOHandler); // a url or ioHandler instance // You can load the dictionary using any api available to the platform const dict loadDictionary(path/to/dict.txt); const model new automl.ImageClassificationModel(graphModel, dict);对应实现见 src/img_classification.tsImageClassificationModel的构造函数就是constructor(public graphModel: GraphModel, public dictionary: string[])没有任何额外初始化逻辑模型与字典只是被持有供后续classify使用。预测classify 与 centerCroptfjs-automl 会自动完成所有图像预处理normalize、resize、crop。你传入的img可以是HTMLImageElement、HTMLCanvasElement、HTMLVideoElement、ImageData或一个 3D 的tf.Tensor。这些类型在源码中被收敛为联合类型ImageInputsrc/types.tsexport type ImageInput ImageData|HTMLImageElement|HTMLCanvasElement|HTMLVideoElement|Tensor3D;预测示例img idimg srcPATH_TO_IMAGE /const img document.getElementById(img); const options {centerCrop: true}; const predictions await model.classify(img, options);options可选当前只有一个属性centerCrop—— 默认true。由于模型期望正方形输入需要先缩放为true时会先对图像做中心裁剪cropped to the center再 resize。返回值predictions是按概率降序排列的标签及概率列表[ {label: daisy, prob: 0.931}, {label: dandelion, prob: 0.027}, {label: roses, prob: 0.013}, ... ]结合源码可以进一步理解这条链路src/img_classification.ts输入统一经imageToTensor处理若是Tensor直接使用否则走browser.fromPixels(img)转成Tensor3Dsrc/util.ts。预处理的目标尺寸固定为IMG_SIZE [224, 224]centerCrop: true时centerCropAndResize取图像的短边计算居中裁剪框归一化坐标[top, left, bottom, right]再调用image.cropAndResize一步完成裁剪与缩放到 224×224centerCrop: false时直接image.resizeBilinear缩放到 224×224不裁剪比例会被拉伸。归一化到[-1, 1]区间(pixels / 127.5) - 1对应常量DIV_FACTOR 127.5、SUB_FACTOR 1。最后调用this.graphModel.predict(preprocessedImg)得到逐类分数与dictionary按下标逐一对应组装为{label, prob}数组并在流程中dispose()中间张量——src/img_classification_test.ts 中有专门的测试断言预测前后tf.memory().numTensors不变验证无内存泄漏。进阶用法直接访问 GraphModel进阶用户可通过model.graphModel访问底层的GraphModel从而调用predict()、execute()、executeAsync()等低层方法获取张量级输出model.dictionary则提供有序的标签列表。例如想自己控制预处理或查看原始 logits 时即可绕过classify的封装直接驱动 graph model。目标检测目标检测模型与分类模型一样输出model.jsondict.txt 若干*.bin同样要求它们可被 Web 端作为静态资源访问。运行官方 Demo目标检测 Demo 位于 demo/object_detectioncd demo/object_detection yarn yarn watch同样会在本地1234端口启动 HTTP 服务器托管 Demo。加载模型import * as automl from tensorflow/tfjs-automl; const modelUrl model.json; // URL to the model.json file. const model await automl.loadObjectDetection(modelUrl);非浏览器平台的分步构造方式与图像分类对称import * as automl from tensorflow/tfjs-automl; import * as tf from tensorflow/tfjs; // You can load the graph model using any IO handler const graphModel await tf.loadGraphModel(string|io.IOHandler); // a url or ioHandler instance // You can load the dictionary using any api available to the platform const dict readDictionary(path/to/dict.txt); const model new automl.ObjectDetectionModel(graphModel, dict);预测detect 与三个选项图像输入类型与分类相同ImageInput联合类型库自动处理 normalize、resize、crop。示例img idimg srcPATH_TO_IMAGE /const img document.getElementById(img); const options {score: 0.5, iou: 0.5, topk: 20}; const predictions await model.detect(img, options);options可选三个参数默认值由 src/object_detection.ts 中sanitizeOptions补齐参数默认值说明score0.5置信度阈值0~1。低于该分数的候选框会被丢弃iou0.5交并比Intersection over Union阈值0~1度量两个框的重叠程度最终输出的框两两重叠不会超过该阈值topk20最多返回置信度最高的topk个目标实际数量可能少于该值返回值是按分数排序的预测对象列表[ { box: { left: 105.1, top: 22.2, width: 70.6, height: 55.7 }, label: Tomato, score: 0.972 }, ... ]Box的坐标语义在 src/object_detection.ts 的类型定义中有精确注释top为距图像顶部的像素数left为距图像左侧的像素数width/height为框的宽高。detect 的源码级执行链路detect的内部流程比classify更有信息量完整阅读 src/object_detection.ts 可以得到以下实现事实喂入固定输入节点图像经cast(expandDims(imageToTensor(input)), float32)后写入以ToFloat为键的 feed 字典INPUT_NODE_NAME ToFloat说明 AutoML Edge 检测模型的图输入节点被固定命名为ToFloat。指定两个输出节点通过graphModel.executeAsync(feedDict, OUTPUT_NODE_NAMES)只取Postprocessor/convert_scores每候选框的类别分数与Postprocessor/Decode/transpose_1解码后的候选框两个张量。逐框取最可能类别calculateMostLikelyLabels将形状为[1, numBoxes, numClasses]的分数张量展平遍历对每个候选框取分数最大的类别下标得到boxScores与boxLabels。非极大值抑制NMS调用image.nonMaxSuppressionAsync(boxesTensor, boxScores, topk, iou, score)一次性完成score阈值过滤、iou去重与topk截断返回选中框的下标数组。坐标还原buildDetectedObjects中模型输出的框坐标是归一化的[top, left, bottom, right]0~1最终转换为像素坐标left * width、top * height、(right - left) * width、(bottom - top) * height其中width/height取自输入图像的宽高——因此返回的box是相对输入图像的像素坐标。整段流程在Promise.all取回两份Float32Array数据后全部在 JS 侧完成并在结束前dispose([img, scoresTensor, boxesTensor, selectedBoxesTensor])释放张量。进阶用法与分类模型相同model.graphModel暴露底层GraphModel可调用predict()、execute()、executeAsync()直接获得张量例如自行实现 NMS 或保留全部候选框model.dictionary提供有序标签列表用于将类别下标映射为标签文本。输入类型、测试与适用前提小结输入统一抽象两类任务都接受ImageInputImageData | HTMLImageElement | HTMLCanvasElement | HTMLVideoElement | Tensor3Dsrc/types.ts非 Tensor 输入统一由browser.fromPixels转换src/util.ts。内存管理可验证分类测试显式断言预测前后张量计数不变src/img_classification_test.ts字典加载的 URL 拼接行为在 src/util_test.ts 中被相对/绝对路径用例完整覆盖。运行环境浏览器环境需要 WebGL 后端参与推理peer dependency非浏览器环境可通过tf.loadGraphModel的任意 IO handler 加载模型并直接构造模型实例这正对应文档中particularly relevant for non-browser platforms的建议。限制说明分类预处理固定为 224×224 与[-1, 1]归一化检测模型的输入/输出节点名ToFloat、Postprocessor/*与 AutoML Edge 导出的图结构绑定因此这两组 API 面向的是 AutoML Edge 产出的标准模型文件不适用于任意图结构的自训模型。综上tfjs-automl 的价值在于以最小的 API 面两个加载函数 两个模型类 一套共享输入类型覆盖了 AutoML Edge 模型的完整推理闭环文档给出的每个参数、每个默认值与返回结构都能在tfjs-automl/src的源码与测试中找到一一对应的实现依据。【免费下载链接】tfjsA WebGL accelerated JavaScript library for training and deploying ML models.项目地址: https://gitcode.com/gh_mirrors/tf/tfjs创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考