TensorFlow.js Op 模块化改造完整指南:以 SquaredDifference 为例的 Kernel/Gradient 迁移实战
人工智能机器学习深度学习前端后端【免费下载链接】tfjsA WebGL accelerated JavaScript library for training and deploying ML models.项目地址https://gitcode.com/gh_mirrors/tf/tfjs点击查看免费下载本文基于 tfjs-core/development/op_modularization.md系统讲解 TensorFlow.js 将 Op 改造为模块化架构的完整工作流——涵盖 Kernel 注册、Op 拆分、链式 API、梯度注册的每一步并结合SquaredDifference的实际源码tfjs-core 与 tfjs-backend-cpu给出可直接对照的落地范例。读完本文你将掌握在 TensorFlow.js 中把一个后端相关的 Op 逐步改造成Op Kernel Gradient三段式模块化结构的方法并理解runKernelFunc与runKernel的演进关系为提交社区 PR 或维护自有分支做好准备。注意文中描述的是 tfjs 模块化改造所追求的目标世界the world we are moving towards部分描述与当前仓库状态存在差异例如 ops/squared_difference.ts 已经改用ENGINE.runKernel而非文档中的runKernelFunc。本文会结合仓库现状逐一标注。一、为什么需要 Op 模块化三个核心概念在动手之前先厘清 TensorFlow.js 中三个容易混淆的核心概念Glossary它们也是整个模块化改造的分层依据Op算子后端无关backend agnostic的函数通常以公开 API 的形式暴露给最终用户实现在tfjs-core中。例如tf.squaredDifference(a, b)以及链式写法a.squaredDifference(b)。Kernel内核针对特定后端的底层实现被一个或多个 Op 复用。Kernel 及其接口定义在tfjs-core/src/kernel_names.ts。Kernel 不能调用其他 Kernel也不能回调用 tfjs 的公开 APIKernel 之间可以通过普通函数导入共享代码。Gradient梯度某个 Kernel 的反向模式backward mode定义同样实现在 tfjs-core 中且是后端无关的——即它们调用其他 Op 或 Kernel 来完成求导。这三者之间的调度桥梁是runKernelFunc它是 tfjs-core engine 中负责执行函数的入口既能处理模块化 Kernel也能处理非模块化 Kernel通过 backend 对象而非 kernel registry 调用。文档明确指出当所有后端的所有 Kernel 都完成模块化后runKernelFunc将被runKernel取代。从当前仓库源码看这个演进已经基本完成engine.ts 中runKernelFunc作为私有方法存在而模块化 op 已统一走runKernel如 ops/squared_difference.ts 中ENGINE.runKernel(SquaredDifference, inputs, attrs)的调用方式。二、改造策略先全部 Op再逐个 Kernel模块化改造有一个明确的总体顺序在开始模块化任何后端的 Kernel 之前先把所有 Op 模块化We will be modularising all the ops before modularizing any of the kernels。这是因为 Op 模块化只涉及 tfjs-core 内部的接口拆分而 Kernel 模块化会同时牵动各后端CPU、WebGL、WASM 等的实现。正式开始前还有一步社区协作要求前往 tfjs 仓库的 issue #2822在评论区告知你要改造哪个 Op避免与他人重复劳动。三、tfjs-core 内的六步改造流程步骤 1在 kernel_names.ts 中添加 Kernel 名称与接口在 tfjs-core/src/kernel_names.ts 中为 Kernel 添加标识符并可选地定义Inputs与Attrs类型export const SquaredDifference SquaredDifference; export type SquaredDifferenceInputs PickNamedTensorInfoMap, a|b;要点命名尽量对齐 TensorFlow 的 C API这是文档中引用的外部参考仓库内不包含该 API 源码无法完全对齐时需向维护者寻求指导。Inputs类型通过PickNamedTensorInfoMap, ...声明NamedTensorInfoMap定义在 kernel_registry.ts。带属性如 axis、keepDims的 Kernel 还需声明Attrs接口参考同文件中的AvgPoolAttrs、AvgPool3DAttrs等示例。步骤 2创建 src/ops/op_name.ts把 Op 定义迁移过来在tfjs-core/src/ops/下新建以 op 命名的文件例如squared_difference.ts。文档给出了一个使用runKernelFunc的过渡期完整实现其中内嵌了前向函数与梯度定义并标注了模块化梯度完成后需要删除的区间Modularization noteimport {ENGINE, ForwardFunc} from ../engine; import {SquaredDifference, SquaredDifferenceInputs} from ../kernel_names; import {Tensor} from ../tensor; import {NamedTensorMap} from ../tensor_types; import {makeTypesMatch} from ../tensor_util; import {convertToTensor} from ../tensor_util_env; import {TensorLike} from ../types; import {assertAndGetBroadcastShape} from ./broadcast_util; import {op} from ./operation; import {scalar} from ./tensor_ops; function squaredDifference_T extends Tensor( a: Tensor|TensorLike, b: Tensor|TensorLike): T { let $a convertToTensor(a, a, squaredDifference); let $b convertToTensor(b, b, squaredDifference); [$a, $b] makeTypesMatch($a, $b); assertAndGetBroadcastShape($a.shape, $b.shape); // **************** // Modularization note: this gradient definition should be removed from // here once the modular gradient is implemented in the steps below. //***************** const der (dy: Tensor, saved: Tensor[]) { const [$a, $b] saved; const two scalar(2); const derA () dy.mul($a.sub($b).mul(two)); const derB () dy.mul($b.sub($a).mul(two)); return {a: derA, b: derB}; }; // **************** // END Modularization note //***************** const forward: ForwardFuncTensor (backend, save) { const res backend.squaredDifference($a, $b); save([$a, $b]); return res; }; const inputs: SquaredDifferenceInputs {a: $a, b: $b}; const attrs {}; const inputsToSave [$a, $b]; const outputToSave: boolean[] []; return ENGINE.runKernelFunc( forward, inputs as unknown as NamedTensorMap, der, SquaredDifference, attrs, inputsToSave, outputToSave) as T; } export const squaredDifference op({squaredDifference_});关键设计原则Op 只做输入校验和数据转换让参数与 Kernel 接口完全匹配其余数据变换一律交给 Kernel。核心准则是Kernel 接口定义的工作不应被拆散到 Op 与 Kernel 之间。由于部分后端如 wasm的 Kernel 已模块化把数据操作从 Op 移交给 Kernel 时旧模块化 Kernel 可能会因此破坏测试失败此时需要同步修改这些 Kernel 以匹配新输入。仍然使用runKernelFunc是为了兼容尚未模块化 Kernel 的后端——这是文档写作时的过渡状态。仓库现状对照当前 ops/squared_difference.ts 已经完成最终形态——梯度定义被移除、前向函数改为直接调用ENGINE.runKernel(SquaredDifference, inputs, attrs)完整体现了文档中过渡态 → 终态的演进路径。步骤 3从 src/ops/ops.ts 导出模块化 Op在 tfjs-core/src/ops/ops.ts 中集中导出所有模块化 Op文件头部注释即标明 Modularized ops.export {squaredDifference} from ./squared_difference;该文件目前已有 343 行、导出上百个模块化 Op如abs、add、conv2d等是公开 API 的汇总出口。步骤 4创建链式 APIchained opaugmentor在src/public/chained_ops/op_name.ts中新建链式方法让Tensor实例可以直接调用import {squaredDifference} from ../../ops/squared_difference; import {Tensor} from ../../tensor; import {Rank, TensorLike} from ../../types; declare module ../../tensor { interface TensorR extends Rank Rank { squaredDifferenceT extends Tensor(b: Tensor|TensorLike): T; } } Tensor.prototype.squaredDifference functionT extends Tensor(b: Tensor| TensorLike): T { this.throwIfDisposed(); return squaredDifference(this, b); };要点通过declare module扩展Tensor接口类型然后挂载原型方法方法内部先调用throwIfDisposed()校验张量未被释放再委托给步骤 2 的模块化 Op。必须把 augmentor 注册到src/public/chained_ops/register_all_chained_ops.ts当前仓库中import ./squared_difference;位于第 135 行并在register_all_chained_ops_test.ts中补充对应的链式调用测试。完成以上步骤后从src/tensor.ts中移除该 Op既要从Tensor类中删除对应方法也要从OpHandler接口中删除保证链式 API 只保留单一实现来源。步骤 5为没有模块化梯度的 Kernel 创建梯度在src/gradients/下按 Kernel 名创建梯度文件例如SquaredDifference_grad.ts。注意梯度中必须使用直接导入的 Op避免使用链式 APIimport {SquaredDifference} from ../kernel_names; import {GradConfig} from ../kernel_registry; import {mul, sub} from ../ops/binary_ops; import {scalar} from ../ops/tensor_ops; import {Tensor} from ../tensor; export const squaredDifferenceGradConfig: GradConfig { kernelName: SquaredDifference, gradFunc: (dy: Tensor, saved: Tensor[]) { const [a, b] saved; const two scalar(2); const derA () mul(dy, mul(two, sub(a, b))); const derB () mul(dy, mul(two, sub(b, a))); return {a: derA, b: derB}; } };从源码结构可以印证GradConfig接口定义在 kernel_registry.ts除kernelName与gradFunc外还支持inputsToSave本次梯度需要保存的输入名、saveAllInputs、outputsToSave。仓库中的最终版本 SquaredDifference_grad.ts 正是通过inputsToSave: [a, b]显式声明需要保存的输入gradFunc再从中取出a、b计算d(a-b)²的两个偏导。步骤 6把梯度配置注册到 register_all_gradients.ts在 tfjs-core/src/register_all_gradients.ts 中导入新梯度配置并加入gradConfigs列表import {squaredDifferenceGradConfig} from ./gradients/SquaredDifference_grad; const gradConfigs: GradConfig[] [ // add the gradient config to this list. squaredDifferenceGradConfig, ];当前仓库中该文件已包含上百个梯度配置如addGradConfig、conv2DGradConfig等并在第 230 行起通过for (const gradientConfig of gradConfigs)循环调用registerGradient完成全局注册供反向传播查询使用。四、提交 PR质量收尾完成以上步骤后即可提交 PR 供审查。提交前务必在tfjs-core目录下本地运行yarn test确保所有单元测试通过——这既是文档明确的强制要求也是模块化改造中防止破坏其他后端 Kernel 的关键防线。五、模块化完成后的全景仓库源码印证结合仓库当前状态可以看到模块化改造的目标形态已经落地1. Kernel 注册机制后端 Kernel 通过 kernel_registry.ts 的registerKernel(config)注册到全局kernelRegistryKernelConfig包含kernelName、backendName、kernelFunc及可选的setupFunc/disposeFunc。查询时通过getKernel(kernelName, backendName)按 kernel_backend 复合键查找。2. 后端 Kernel 实现例如 CPU 后端的 SquaredDifference.ts 通过binaryKernelFunc与createSimpleBinaryKernelImpl实现逐元素(a-b)²计算并导出squaredDifferenceConfig供注册——这正是文档所说Kernel 之间通过普通函数导入共享代码的具体体现。3. 调度入口统一模块化 Op 不再直接调用backend.squaredDifference(...)而是把SquaredDifference名称、inputs、attrs 交给ENGINE.runKernel由 engine 在运行时根据当前激活的 backend 从 kernel registry 中查找对应实现。整个链路Op 定义 → 链式 API → Kernel 注册 → 梯度注册完全由kernel_names.ts中的常量串起来实现了名称即契约。六、改造清单速查Checklist阶段文件动作接口tfjs-core/src/kernel_names.ts添加 Kernel 常量、Inputs/Attrs 类型Op 定义tfjs-core/src/ops/squared_difference.ts迁移 Op仅做校验与接口匹配Op 导出tfjs-core/src/ops/ops.ts添加export {...} from ./xxx链式 APItfjs-core/src/public/chained_ops/squared_difference.ts声明 module 扩展 原型方法链式注册/测试tfjs-core/src/public/chained_ops/register_all_chained_ops.ts注册 augmentor、补充测试清理旧 APItfjs-core/src/tensor.ts移除 Tensor 类方法与 OpHandler 接口项梯度tfjs-core/src/gradients/SquaredDifference_grad.ts创建 GradConfig使用直接导入的 Op梯度注册tfjs-core/src/register_all_gradients.ts加入 gradConfigs 列表验证tfjs-core目录下运行yarn test按照这个清单逐项推进你就能以最小的破坏性完成一个 Op 的模块化改造从Op 内嵌 backend 调用与梯度定义的过渡形态演进为Op校验→ Kernel后端实现→ Gradient反向传播职责清晰、后端无关的模块化架构为后续各后端 Kernel 的逐一模块化铺平道路。赞分享人工智能机器学习深度学习前端后端【免费下载链接】tfjsA WebGL accelerated JavaScript library for training and deploying ML models.项目地址https://gitcode.com/gh_mirrors/tf/tfjs点击查看免费下载相关推荐ESLint v9 Flat Config 迁移实战指南以 Chainlit Monorepo 的 .eslintrc 到 eslint.config.mjs 改造为例ESLint v9 Flat Config 迁移实战指南以 Chainlit Monorepo 的 .eslintrc 到 eslint.config.mjs人工智能大模型AI 应用后端前端Bazel 从 Maven 迁移指南以 Guava 项目为例的完整实战教程Bazel 从 Maven 迁移指南以 Guava 项目为例的完整实战教程 本指南基于 Bazel 官方文档 docs/migrate/maven.mdx构建工具Starship v0.45.0 迁移实战指南prompt_order 与模块 prefix/suffix 的统一 format 化改造Starship v0.45.0 迁移实战指南prompt_order 与模块 prefix/suffix 的统一 format 化改造 Starship 在CLI开发工具创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →