尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

brain.js 实战指南:在浏览器与 Node.js 中构建 GPU 加速的 JavaScript 神经网络

brain.js 实战指南:在浏览器与 Node.js 中构建 GPU 加速的 JavaScript 神经网络 人工智能机器学习深度学习【免费下载链接】brain.js GPU accelerated Neural networks in JavaScript for Browsers and Node.js项目地址https://gitcode.com/gh_mirrors/br/brain.js点击查看免费下载导读brain.js 是一个用 JavaScript 编写的 GPU 加速神经网络库可同时运行于浏览器与 Node.js 环境。本文以仓库 README.md 为主线完整梳理从安装、训练数据格式、训练选项到train()/run()/forecast()等核心方法的使用方式并结合仓库源码如 src/neural-network.ts、src/recurrent/rnn.ts讲解各参数在底层实现中的作用。阅读完本文你将能够独立完成前馈网络Feedforward、时间步网络TimeStep与循环网络RNN/LSTM/GRU、自编码器AE的搭建、训练、序列化与部署。一、安装与使用通过 NPM 安装在 Node.js 项目中执行npm install brain.js安装后通过require或import引入。仓库package.json中main指向dist/index.js浏览器端构建产物为dist/browser.jsunpkg字段TypeScript 类型声明位于dist/目录。通过 CDN 引入在 HTML 中直接引入浏览器构建版本script src//unpkg.com/brain.js/script也可以直接下载最新的浏览器构建文件使用。安装注意事项与源码构建brain.js 的 GPU 支持依赖原生模块headless-gl。大多数情况下npm install brain.js即可直接工作如果安装失败通常是预编译二进制无法从 GitHub 下载需要自行从源码构建。构建前请确保系统依赖已安装然后运行npm rebuild各平台的系统依赖如下Mac OS X受支持的 Python 版本、XCode。Ubuntu/Debian受支持的 Python 版本、GNU C 环境build-essential、libxi-dev、可用的 OpenGL 驱动、GLEW、pkg-config。可用一条命令安装sudo apt-get install -y build-essential libglew-dev libglu1-mesa-dev libxi-dev pkg-configWindows受支持的 Python 版本、Microsoft Visual Studio Build Tools 2022老版本 npm 可配合npm config set msvs_version 2022与npm config set python python3使用。从源码构建的完整脚本定义见 package.json 的scripts字段build会并行执行 browser 与 node 的 rollup 构建及 TypeScript 声明生成。二、快速上手XOR 示例前馈神经网络逼近 XOR// provide optional config object (or undefined). Defaults shown. const config { binaryThresh: 0.5, hiddenLayers: [3], // array of ints for the sizes of the hidden layers in the network activation: sigmoid, // supported activation types: [sigmoid, relu, leaky-relu, tanh], leakyReluAlpha: 0.01, // supported for activation type leaky-relu }; // create a simple feed-forward neural network with backpropagation const net new brain.NeuralNetwork(config); net.train([ { input: [0, 0], output: [0] }, { input: [0, 1], output: [1] }, { input: [1, 0], output: [1] }, { input: [1, 1], output: [0] }, ]); const output net.run([1, 0]); // [0.987]其中binaryThresh是二值化阈值默认0.5hiddenLayers指定隐藏层结构activation指定激活函数leakyReluAlpha仅在激活函数为leaky-relu时生效。这些默认值在 src/neural-network.ts 的defaults()与trainDefaults()中有对应实现。循环神经网络逼近 XOR// provide optional config object, defaults shown. const config { inputSize: 20, inputRange: 20, hiddenLayers: [20, 20], outputSize: 20, learningRate: 0.01, decayRate: 0.999, }; // create a simple recurrent neural network const net new brain.recurrent.RNN(config); net.train([ { input: [0, 0], output: [0] }, { input: [0, 1], output: [1] }, { input: [1, 0], output: [1] }, { input: [1, 1], output: [0] }, ]); let output net.run([0, 0]); // [0] output net.run([0, 1]); // [1] output net.run([1, 0]); // [1] output net.run([1, 1]); // [0]循环网络的配置还包含inputRange输入取值范围与decayRate学习率衰减系数其默认值可在 src/recurrent/rnn.ts 中查看。需要注意用神经网络求解 XOR 只是一个教学演示对于真实的分类/回归任务应使用贴近业务的数据集。三、训练Training调用train()时网络必须一次性批量接收全部训练数据。训练模式越多耗时通常越长但网络对未见过的模式分类能力通常也越强。3.1 训练数据格式不同网络类型要求不同的数据格式仓库为每种类型都配有对应的单元测试与端到端测试见 src/ 下各.test.ts文件。面向NeuralNetwork的数据每条训练样本应包含input与output二者既可以是0到1之间的数字数组也可以是取值范围为0到1的数字哈希hash。以颜色对比度识别为例const net new brain.NeuralNetwork(); net.train([ { input: { r: 0.03, g: 0.7, b: 0.5 }, output: { black: 1 } }, { input: { r: 0.16, g: 0.09, b: 0.2 }, output: { white: 1 } }, { input: { r: 0.5, g: 0.5, b: 1.0 }, output: { white: 1 } }, ]); const output net.run({ r: 1, g: 0.4, b: 0 }); // { white: 0.99, black: 0.002 }值得注意的是不同样本的输入对象结构不必一致。brain.js 会通过 lookup 表为所有出现过的键分配统一的输入/输出维度input或output中未出现的键按 0 处理net.train([ { input: { r: 0.03, g: 0.7 }, output: { black: 1 } }, { input: { r: 0.16, b: 0.2 }, output: { white: 1 } }, { input: { r: 0.5, g: 0.5, b: 1.0 }, output: { white: 1 } }, ]); const output net.run({ r: 1, g: 0.4, b: 0 }); // { white: 0.81, black: 0.18 }上述 lookup 机制在 src/lookup.ts 与 src/neural-network.ts 的getTypedArrayFn中实现哈希输入会被转换为定长的Float32Array缺失的键以 0 填充。面向RNNTimeStep、LSTMTimeStep与GRUTimeStep的数据每条训练样本可以是一个数字数组一个由数字数组组成的数组二维数组。使用数字数组时间序列预测下一个值const net new brain.recurrent.LSTMTimeStep(); net.train([[1, 2, 3]]); const output net.run([1, 2]); // 3使用二维数组预测下一组向量const net new brain.recurrent.LSTMTimeStep({ inputSize: 2, hiddenLayers: [10], outputSize: 2, }); net.train([ [1, 3], [2, 2], [3, 1], ]); const output net.run([ [1, 3], [2, 2], ]); // [3, 1]面向RNN、LSTM与GRU的数据每条训练样本可以是一个值数组一个字符串带input与output的对象其中input/output可以是值数组或字符串。注意使用值数组时任意数值都可以作为输入但每个不同的值在网络中都由单个输入神经元表示——因此不同的值越多输入层就越大。如果数据是成百上千甚至数百万个浮点值这类网络并不适合此外非字符串数据的支持仍处于 beta 阶段。直接使用字符串训练文字续写const net new brain.recurrent.LSTM(); net.train([ doe, a deer, a female deer, ray, a drop of golden sun, me, a name I call myself, ]); const output net.run(doe); // , a deer, a female deer使用带input/output的字符串做情感分类const net new brain.recurrent.LSTM(); net.train([ { input: I feel great about the world!, output: happy }, { input: The world is a terrible place!, output: sad }, ]); const output net.run(I feel great about the world!); // happy字符串数据会由 src/utilities/data-formatter.ts 中的DataFormatter编码为字符索引序列RNN的run()再根据maxPredictionLength与温度参数逐字符采样生成结果。面向AE自编码器的数据每条训练样本可以是数字数组或数字数组的数组。用自编码器压缩 XOR 计算结果const net new brain.AE( { hiddenLayers: [ 5, 2, 5 ] } ); net.train([ [ 0, 0, 0 ], [ 0, 1, 1 ], [ 1, 0, 1 ], [ 1, 1, 0 ] ]);编码/解码const input [ 0, 1, 1 ]; const encoded net.encode(input); const decoded net.decode(encoded);去噪const noisyData [ 0, 1, 0 ]; const data net.denoise(noisyData);异常检测const shouldBeFalse net.includesAnomalies([0, 1, 1]); const shouldBeTrue net.includesAnomalies([0, 1, 0]);AE的encode/decode/denoise/includesAnomalies在 src/autoencoder.ts 中实现encode与decode在对应子网络未训练时会抛出UntrainedNeuralNetworkError见 src/errors/untrained-neural-network-error.tsincludesAnomalies通过比较输入与去噪重建结果的差异向量判断是否存在异常。3.2 训练选项Training Optionstrain()的第二个参数是训练选项哈希下表完整列出各选项及其默认值与校验规则校验逻辑可参考 src/neural-network.trainopts.test.tsnet.train(data, { // Defaults values -- expected validation iterations: 20000, // the maximum times to iterate the training data -- number greater than 0 errorThresh: 0.005, // the acceptable error percentage from training data -- number between 0 and 1 log: false, // true to use console.log, when a function is supplied it is used -- Either true or a function logPeriod: 10, // iterations between logging out -- number greater than 0 learningRate: 0.3, // scales with delta to effect training rate -- number between 0 and 1 momentum: 0.1, // scales with next layers change value -- number between 0 and 1 callback: null, // a periodic call back that can be triggered while training -- null or function callbackPeriod: 10, // the number of iterations through the training data between callback calls -- number greater than 0 timeout: number, // the max number of milliseconds to train for -- number greater than 0. Default -- Infinity });网络在以下两个条件之一满足时停止训练训练误差低于阈值默认0.005或迭代次数达到上限默认20000。log / logPeriod默认训练全程不输出信息设置log: true会按logPeriod间隔向控制台打印当前训练误差误差应持续下降。将log设置为函数时该函数会接收更新数据而不是打印。callback / callbackPeriod若想在自有输出中使用训练进度值可将callback设为函数它每隔callbackPeriod次迭代被调用一次。learningRate0到1之间的数值影响训练速度。接近0训练更慢但更稳接近1训练更快但结果可能收敛到局部最小值在新数据上表现变差过拟合。默认0.3。momentum与学习率类似取值0到1乘以下一层的变化量来平滑更新。默认0.1。timeout训练的最长毫秒数默认Infinity。源码中通过Date.now() this.trainOpts.timeout计算截止时间并在训练循环中检查。以上训练选项既可以传入构造函数也可以通过updateTrainingOptions(opts)方法在之后更新它们会保存在网络上并在训练时生效。将网络保存为 JSON 时训练选项会一并保存与恢复——唯一例外是callback恢复后被丢弃与log恢复后使用console.log。这一行为对应 src/neural-network.ts 中updateTrainingOptions与fromJSON的实现。此外网络有一个默认值为true的布尔属性invalidTrainOptsShouldThrow当它为true时传入超出正常范围的训练选项会抛出异常并附带说明设为false时不再抛错但仍会向console.warn输出相关提示。3.3 异步训练Async TrainingtrainAsync()的参数与train()相同数据与选项但它不直接返回训练结果对象而是返回一个 Promiseresolve 时得到训练结果。不支持trainAsync()的类brain.recurrent.RNNbrain.recurrent.GRUbrain.recurrent.LSTMbrain.recurrent.RNNTimeStepbrain.recurrent.GRUTimeStepbrain.recurrent.LSTMTimeStep基本用法const net new brain.NeuralNetwork(); net .trainAsync(data, options) .then((res) { // do something with my trained network }) .catch(handleError);多网络并行训练const net new brain.NeuralNetwork(); const net2 new brain.NeuralNetwork(); const p1 net.trainAsync(data, options); const p2 net2.trainAsync(data, options); Promise.all([p1, p2]) .then((values) { const res values[0]; const res2 values[1]; console.log( net trained in ${res.iterations} and net2 trained in ${res2.iterations} ); // do something super cool with my 2 trained networks }) .catch(handleError);3.4 交叉验证Cross Validation对于较大数据集交叉验证可以提供更稳健的训练方式。brain.js 的 API 用法如下const crossValidate new brain.CrossValidate(() new brain.NeuralNetwork(networkOptions)); crossValidate.train(data, trainingOptions, k); //note k (or KFolds) is optional const json crossValidate.toJSON(); // all stats in json as well as neural networks const net crossValidate.toNeuralNetwork(); // get top performing net out of crossValidate // optionally later const json crossValidate.toJSON(); const net crossValidate.fromJSON(json);CrossValidate支持以下类brain.NeuralNetworkbrain.RNNTimeStepbrain.LSTMTimeStepbrain.GRUTimeStep从源码看src/cross-validate.ts 会按折k-fold划分训练/测试集汇总每折的trainTime、testTime、iterations、error并在二分类场景下额外统计truePos/trueNeg/falsePos/falseNeg及precision、recall、accuracy最终通过toNeuralNetwork()返回表现最优的网络。四、核心方法Methodstrain(trainingData)- trainingStatustrain()返回一个描述训练过程的哈希{ error: 0.0039139985510105032, // training error iterations: 406 // training iterations }run(input)- predictionrun()支持的类brain.NeuralNetworkbrain.NeuralNetworkGPU具备brain.NeuralNetwork的全部功能但运行在 GPU 上通过 gpu.js 使用 WebGL2、WebGL1或回退到 CPUbrain.recurrent.RNNbrain.recurrent.LSTMbrain.recurrent.GRUbrain.recurrent.RNNTimeStepbrain.recurrent.LSTMTimeStepbrain.recurrent.GRUTimeStep示例// feed forward const net new brain.NeuralNetwork(); net.fromJSON(json); net.run(input); // time step const net new brain.LSTMTimeStep(); net.fromJSON(json); net.run(input); // recurrent const net new brain.LSTM(); net.fromJSON(json); net.run(input);NeuralNetworkGPU在 src/neural-network-gpu.ts 中实现其构造函数接收{ mode: cpu | gpu }选项并通过new GPU({ mode })创建计算上下文gpu.js是它的 peerDependency见 package.json。forecast(input, count)- predictionsforecast()可用于以下类输出一个预测数组即对输入序列的延续brain.recurrent.RNNTimeStepbrain.recurrent.LSTMTimeStepbrain.recurrent.GRUTimeStep示例const net new brain.LSTMTimeStep(); net.fromJSON(json); net.forecast(input, 3);toJSON() - json与fromJSON(json)toJSON()将神经网络序列化为 JSONfromJSON(json)从 JSON 反序列化恢复网络。五、训练失败排查Failing如果网络训练失败最终误差会高于误差阈值。可能的原因通常是训练数据噪声太大最常见、网络的隐藏层或节点数不足以处理数据复杂度、或迭代次数不够。若训练 20000 次迭代后误差仍高达0.4左右说明网络基本无法理解给定数据。RNN / LSTM / GRU 输出过长或过短网络的maxPredictionLength属性默认100用于调节输出长度。例如在用若干小说训练后想续写一部新小说const net new brain.recurrent.LSTM(); // later in code, after training on a few novels, write me a new one! net.maxPredictionLength 1000000000; // Be careful! net.run(Once upon a time);该默认值定义在 src/recurrent/rnn.ts 的defaults中run()实际以maxPredictionLength 输入长度作为最大生成步数并在此上限内按采样逻辑逐字符生成超限即停止。六、JSON 序列化将训练好的网络状态保存为 JSON 或从 JSON 加载const json net.toJSON(); net.fromJSON(json);配合toJSON()可以在离线环境或 Web Worker中完成昂贵的训练再把训练好的网络通过 JSON 部署到网页上。七、独立函数Standalone Function可以从训练好的网络得到一个行为与run()完全一致的独立函数const run net.toFunction(); const output run({ r: 1, g: 0.4, b: 0 }); console.log(run.toString()); // copy and paste! no need to import brain.jsrun.toString()会输出该函数的完整源码可以直接复制粘贴使用无需再引入 brain.js。八、网络选项OptionsNeuralNetwork()接受一个选项哈希const net new brain.NeuralNetwork({ activation: sigmoid, // activation function hiddenLayers: [4], learningRate: 0.6, // global learning rate, useful when training using streams });activation指定神经网络使用的激活函数当前支持四种sigmoid 为默认值sigmoidreluleaky-relu相关选项leakyReluAlpha可选数字默认0.01tanh仓库中每种激活函数都有独立实现与测试src/activation/sigmoid.ts、src/activation/relu.ts、src/activation/leaky-relu.ts、src/activation/tanh.ts并在 src/activation/index.ts 统一导出。这些激活函数同样以 layer 形式存在于 src/layer/ 目录如 src/layer/sigmoid.ts、src/layer/leaky-relu.ts供自定义网络架构使用。hiddenLayers指定网络的隐藏层数量与每层节点数。例如需要两个隐藏层第一层 3 个节点、第二层 4 个节点hiddenLayers: [3, 4];默认情况下 brain.js 使用一个隐藏层其大小与输入数组大小成比例。九、流式训练Streamsbrain.js 本身不内置流式训练官方推荐使用train-stream包向NeuralNetwork流式传入数据适用于需要持续接收增量训练数据的场景。与之配合时learningRate作为全局学习率在流式训练中特别有用。十、工具函数Utilitieslikely返回网络输出中置信度最高的键const likely require(brain/likely); const key likely(input, net);其实现位于 src/likely.ts调用net.run(input)后遍历输出对象找出数值最大的属性名作为结果若传入的net为空则抛出TypeError。toSVG将前馈网络的拓扑结构渲染为 SVGscript src../../src/utilities/svg.js/scriptdocument.getElementById(result).innerHTML brain.utilities.toSVG( network, options );实现位于 src/utilities/to-svg.ts仓库内也有对应测试 src/utilities/to-svg.test.ts可用于可视化网络结构、辅助调试。十一、神经网络类型总览Neural Network Types类型源码位置说明brain.NeuralNetworksrc/neural-network.ts带反向传播的前馈神经网络brain.NeuralNetworkGPUsrc/neural-network-gpu.ts前馈神经网络GPU 版本brain.AEsrc/autoencoder.ts自编码器支持反向传播与 GPUbrain.recurrent.RNNTimeStepsrc/recurrent/rnn-time-step.ts时间步循环神经网络brain.recurrent.LSTMTimeStepsrc/recurrent/lstm-time-step.ts时间步长短期记忆网络brain.recurrent.GRUTimeStepsrc/recurrent/gru-time-step.ts时间步门控循环单元brain.recurrent.RNNsrc/recurrent/rnn.ts循环神经网络brain.recurrent.LSTMsrc/recurrent/lstm.ts长短期记忆网络brain.recurrent.GRUsrc/recurrent/gru.ts门控循环单元brain.FeedForwardsrc/feed-forward.ts高度可定制的前馈网络带反向传播brain.Recurrentsrc/recurrent.ts高度可定制的循环网络带反向传播所有公开 APIactivation、AE、CrossValidate、likely、layer、praxis、FeedForward、NeuralNetwork、NeuralNetworkGPU、Recurrent、recurrent、utilities等统一从 src/index.ts 导出recurrent命名空间下包含六个时间步/循环类。为什么需要不同类型的神经网络不同类型的网络擅长不同的任务例如前馈神经网络Feedforward擅长对简单事物分类但它没有对之前动作的记忆结果存在无限种变化。时间步循环网络Time Step Recurrent具备记忆可以预测未来值适合时间序列预测。循环神经网络Recurrent具备记忆并且结果集合是有限的适合字符串续写、序列生成等任务。实际选型时应根据任务是否需要记忆、输出是连续值还是有限集合、数据是序列还是独立样本在上述类型中选择合适的网络。十二、部署建议与注意事项训练是计算密集型任务建议在离线环境或 Web Worker 中完成训练再通过toFunction()或toJSON()将预训练网络部署到网站避免阻塞页面主线程。GPU 加速默认安装即可在多数环境获得 GPU 支持若headless-gl原生模块安装失败参考本文第一节的源码构建步骤。数据范围NeuralNetwork的输入/输出需归一化到0到1循环网络对字符串/有限值集合更友好海量浮点值应改用其他方案。选项校验可通过invalidTrainOptsShouldThrow控制非法训练选项是抛错还是仅告警需要运行时调整训练参数时使用updateTrainingOptions(opts)。本仓库为只读镜像安装、运行与配置方式均可直接参考上文更多源码级细节可继续深入 src/ 目录下的实现与对应.test.ts测试文件。赞分享人工智能机器学习深度学习【免费下载链接】brain.js GPU accelerated Neural networks in JavaScript for Browsers and Node.js项目地址https://gitcode.com/gh_mirrors/br/brain.js点击查看免费下载相关推荐Brain.js 终极指南JavaScript中的GPU加速神经网络框架Brain.js 终极指南JavaScript中的GPU加速神经网络框架 Brain.js 是一个强大的GPU加速神经网络库专门为浏览器和Node.js环境人工智能机器学习深度学习DLSS Swapper 使用指南如何快速切换游戏中的 DLSS、FSR 与 XeSS 版本DLSS Swapper 使用指南如何快速切换游戏中的 DLSS、FSR 与 XeSS 版本 DLSS Swapper 是一款免费的 Windows 工具帮桌面应用WebGL1与WebGL2完全指南brain.js GPU加速神经网络浏览器兼容性终极解析WebGL1与WebGL2完全指南brain.js GPU加速神经网络浏览器兼容性终极解析 brain.js作为领先的JavaScript神经网络库通过GP人工智能机器学习深度学习上一篇Screenshot-to-code与版本控制系统集成Git提交自动化方案下一篇Gobuster正则表达式过滤精准匹配目标路径创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表