简介:这是一套面向机器学习初学者与教学者的前端可视化教学工具,聚焦算法原理理解与交互式实践,特别适用于高校课程演示、自学入门及原理讲解场景。资源基于TensorFlow.js实现浏览器端模型训练与推理,结合D3.js构建动态数据图表与交互界面,完整呈现线性回归、KNN分类与决策树三大经典算法的运行过程、参数影响及预测效果。压缩包共50个文件,含8个核心HTML页面(含index.html及各算法独立演示页)、6个JS脚本(封装模型逻辑与D3渲染)、18张PNG/JPG示意图(展示算法流程与结果可视化)、3个JSON配置与示例数据,以及说明文件.txt和附赠资源.docx等辅助文档,整体仅1.08MB,轻量易部署。目前已有56人下载学习,用户可直接上传自定义数据集、实时拖拽调整超参数(如K值、树深度、学习率),即时观察模型变化,真正实现“所见即所得”的算法认知闭环。
1. 这不是“前端画个图就完事”的玩具:一个能真正讲清线性回归斜率怎么动、KNN决策边界怎么跳、决策树分裂点怎么选的交互式教学沙盒
你有没有试过给学生讲「为什么线性回归的损失曲面是碗状的」,结果他们盯着静态PPT上那张三维等高线图,眼神逐渐放空?有没有调试过KNN分类器,却说不清k=3和k=5时决策边界为何在某几个点突然“撕裂”?这不是学生没悟性,而是传统教学工具缺了一层关键能力:让算法参数变成可拖拽的旋钮,让数学公式变成实时变形的图形,让抽象假设(比如“数据服从独立同分布”)在上传一组异常点后立刻显形为红色警告框。这个基于TensorFlow.js + D3.js构建的小程序,正是为解决这类“原理看不见、调参摸不着、错误难归因”的教学痛点而生——它不跑真实业务数据,但每一步推导都严格对应《机器学习》(周志华西瓜书)、吴恩达课程中的数学定义;它不追求大屏炫酷效果,但每个坐标轴刻度、每条拟合线斜率、每个KNN投票圆圈半径,都绑定着可验证的TensorFlow.js张量计算结果。适合高校助教快速搭建算法演示页、培训机构制作原理动画课件、自学入门者亲手拖动滑块理解偏差-方差权衡。核心价值不在“用了什么技术”,而在“所有可视化元素背后,都有可打断、可inspect、可修改的JS代码链路”。
2. 从零搭起可交互沙盒:TensorFlow.js加载模型 + D3.js驱动视图更新的最小闭环
这个小程序的本质,是一个双引擎协同系统:TensorFlow.js负责“算得准”(数值计算、梯度下降、预测推理),D3.js负责“看得懂”(坐标映射、动态过渡、事件绑定)。二者不能简单拼接,必须建立明确的数据契约——即哪些变量由TF.js输出,哪些DOM元素由D3.js控制,中间如何触发更新。下面以线性回归模块为例,拆解最简可行路径。
2.1 初始化TensorFlow.js环境并定义可训练模型
我们不直接用tf.layers.model封装复杂网络,而是手写最基础的单变量线性回归:y = w * x + b。这样做的好处是——所有参数(w, b)全程暴露在JS作用域中,便于D3绑定滑块事件;同时避免Keras层抽象带来的黑匣子感,学生能一眼看懂loss = tf.mean(tf.square(pred.sub(y)))这行代码在算什么。
// linearRegression.js import * as tf from '@tensorflow/tfjs'; export class LinearRegressor { constructor() { // 权重w和偏置b初始化为随机值,但限定范围便于教学观察 this.w = tf.variable(tf.scalar(Math.random() * 2 - 1)); // [-1, 1] this.b = tf.variable(tf.scalar(Math.random() * 2 - 1)); this.learningRate = 0.01; this.optimizer = tf.train.sgd(this.learningRate); } predict(x) { // x是tf.tensor1d([x1, x2, ...]),返回y_pred = w*x + b return this.w.mul(x).add(this.b); } trainStep(x, y) { // 定义损失函数:均方误差 const pred = this.predict(x); const loss = tf.mean(tf.square(pred.sub(y))); // 计算梯度并更新参数 this.optimizer.minimize(() => loss, false, [this.w, this.b]); // 返回当前loss值供D3更新loss曲线 return loss.dataSync()[0]; } getParams() { // 同步获取当前w和b值,供D3渲染直线 return { w: this.w.dataSync()[0], b: this.b.dataSync()[0] }; } }注意:这里
getParams()必须用.dataSync()而非.array(),因为array()返回Promise,会引入异步延迟,导致D3更新滞后于参数变化——这是教学演示中最致命的“不同步”问题。我踩过坑:当学生拖动学习率滑块时,直线跳变比滑块停止晚300ms,直接破坏“参数→效果”的因果直觉。
2.2 用D3.js构建可拖拽的参数控制面板与实时绘图区
D3不负责计算,只做三件事:(1)监听HTML滑块(<input type="range">)的input事件;(2)将滑块值映射为TF.js模型参数;(3)根据TF.js返回的getParams()重绘直线。关键在于事件流设计:滑块改变 → 触发TF.jssetWeights()→ 调用trainStep()单步训练 → 获取新参数 → D3重绘。整个链路必须同步阻塞,否则出现“滑块已停,直线还在动”的玄学现象。
// d3Controller.js import { LinearRegressor } from './linearRegression.js'; const regressor = new LinearRegressor(); let xData = tf.tensor1d([1, 2, 3, 4, 5]); // 示例数据 let yData = tf.tensor1d([2.1, 3.9, 6.2, 7.8, 10.1]); // 绑定学习率滑块 d3.select('#lr-slider') .on('input', function() { const lr = +this.value; regressor.learningRate = lr; // 立即用新学习率执行一次训练步,使直线响应滑块 regressor.trainStep(xData, yData); updatePlot(); // 更新D3视图 }); function updatePlot() { const params = regressor.getParams(); // D3重绘直线:y = params.w * x + params.b const lineGen = d3.line() .x(d => xScale(d)) .y(d => yScale(params.w * d + params.b)); d3.select('#regression-line') .datum(d3.range(0, 10, 0.1)) // 生成x轴采样点 .attr('d', lineGen); }逻辑说明:
updatePlot()里d3.range(0,10,0.1)生成100个x值,代入当前params.w和params.b计算y,再通过D3的line()生成SVG路径。这里不用TF.js张量运算,因为纯JS计算足够快,且避免tensor内存泄漏——TensorFlow.js在频繁创建销毁tensor时,若未手动dispose(),内存占用会指数级增长,页面卡顿。这是血泪经验:曾因忘记xData.dispose(); yData.dispose();,连续拖动10次滑块后内存飙到1.2GB。
2.3 数据上传模块:解析CSV并转换为TensorFlow.js兼容格式
教学场景下,学生常想用自己的数据(如身高体重、房价面积)。前端需支持拖拽上传CSV,并做三件事:(1)校验列数(线性回归要求至少2列);(2)过滤非数字行;(3)归一化处理防止梯度爆炸。特别注意:TensorFlow.js的tensor必须是float32,而CSV解析默认是string,类型错位会导致mul is not a function等静默失败。
// dataUploader.js export function parseCSV(csvText) { const lines = csvText.split('\n').filter(l => l.trim() !== ''); const headers = lines[0].split(',').map(h => h.trim()); if (headers.length < 2) { throw new Error('CSV must have at least 2 columns'); } const data = []; for (let i = 1; i < lines.length; i++) { const values = lines[i].split(',').map(v => parseFloat(v.trim())); if (values.some(isNaN)) continue; // 跳过含非数字行 data.push(values); } if (data.length === 0) throw new Error('No valid numeric data found'); // 提取前两列作为x,y(教学简化) const x = data.map(row => row[0]); const y = data.map(row => row[1]); // 归一化:x = (x - mean) / std,避免学习率失效 const xMean = d3.mean(x); const xStd = d3.deviation(x); const xNorm = x.map(v => (v - xMean) / xStd); return { x: tf.tensor1d(xNorm, 'float32'), // 强制指定dtype y: tf.tensor1d(y, 'float32'), originalX: x, // 保存原始值用于坐标轴标签 originalY: y }; } // 使用示例 document.getElementById('csv-upload').addEventListener('change', async (e) => { const file = e.target.files[0]; const text = await file.text(); try { const tensors = parseCSV(text); xData = tensors.x; yData = tensors.y; // 重置模型参数,避免旧权重干扰新数据 regressor.w.assign(tf.scalar(Math.random() * 0.2 - 0.1)); regressor.b.assign(tf.scalar(Math.random() * 0.2 - 0.1)); updatePlot(); } catch (err) { alert(`数据解析失败: ${err.message}`); } });参数说明:
tf.tensor1d(xNorm, 'float32')中'float32'不可省略。若传入[1,2,3],TF.js默认推断为'int32',后续mul()操作会报错。归一化用d3.mean/deviation而非TF.js内置统计函数,因为D3的标量计算更轻量,且避免在tensor未dispose时创建新tensor。
3. KNN与决策树模块:如何让“距离”和“分裂”在屏幕上肉眼可见
线性回归是连续优化,KNN和决策树则是离散决策。它们的可视化难点在于:KNN的“k值”改变时,决策边界不是平滑变形,而是突变式重组;决策树的“最大深度”调整后,整棵树结构可能完全重构。D3无法像画直线那样简单重绘,必须设计状态管理机制。
3.1 KNN模块:用D3动态渲染投票圆圈与决策热力图
KNN的核心是“找最近的k个邻居”。可视化分两层:(1)实例层:每个样本点旁画一个半径为r的圆,圆内包含其k个最近邻;(2)决策层:在背景网格上用颜色深浅表示该位置被预测为哪一类的概率。关键挑战是——圆半径r必须随k动态缩放,否则k=1时圆太小看不见,k=10时圆覆盖全图。
// knnVisualizer.js export class KNNVisualizer { constructor(data, labels) { this.data = data; // [[x1,y1], [x2,y2], ...] this.labels = labels; // [0,1,0,1,...] this.k = 3; } // 计算点p到所有数据点的欧氏距离,返回k个最近邻索引 getKNearest(p) { const distances = this.data.map((d, i) => ({ idx: i, dist: Math.sqrt((d[0]-p[0])**2 + (d[1]-p[1])**2) })).sort((a,b) => a.dist - b.dist); return distances.slice(0, this.k).map(d => d.idx); } // 渲染单个查询点p的KNN圆圈 renderCircle(p, svgGroup) { const neighbors = this.getKNearest(p); const maxDist = neighbors.length > 0 ? Math.max(...neighbors.map(i => Math.sqrt((this.data[i][0]-p[0])**2 + (this.data[i][1]-p[1])**2))) : 0; // 动态半径:k越大,半径按log缩放,避免重叠 const radius = Math.log(this.k + 1) * 20 + 5; svgGroup.append('circle') .attr('cx', xScale(p[0])) .attr('cy', yScale(p[1])) .attr('r', radius) .attr('fill', 'none') .attr('stroke', '#4A90E2') .attr('stroke-width', 1.5) .attr('stroke-dasharray', '5,5'); // 标出k个邻居点 neighbors.forEach(i => { svgGroup.append('circle') .attr('cx', xScale(this.data[i][0])) .attr('cy', yScale(this.data[i][1])) .attr('r', 4) .attr('fill', colorScale(this.labels[i])); }); } }为什么用
Math.log(k+1)*20+5?实测发现:k=1时半径需≥5px才可见;k=10时若用线性缩放(如k*5),半径=50px会覆盖半个画布。对数缩放让k=1~10时半径落在5~35px合理区间,既保证k=1时清晰,又避免k=10时遮挡。
3.2 决策树模块:用D3.tree生成可折叠的树形图,并绑定节点分裂条件
决策树可视化不是画一棵静态树,而是让用户理解“为什么在这里分裂”。每个内部节点需显示:(1)分裂特征(如“x<3.2”);(2)样本数;(3)基尼不纯度。叶子节点显示预测类别。难点在于——当用户调整“最大深度”时,整棵树结构重建,D3需高效diff旧树与新树。
// treeBuilder.js export function buildDecisionTree(data, labels, maxDepth = 3) { // 递归构建树节点对象(非TF.js tensor,纯JS对象) function splitNode(samples, depth) { if (depth >= maxDepth || samples.length < 2) { return { type: 'leaf', label: majorityVote(labels.filter((_,i) => samples.includes(i))), count: samples.length }; } // 找最佳分裂:遍历所有特征和阈值,选基尼增益最大者 let bestGain = -Infinity; let bestSplit = null; for (let feat = 0; feat < 2; feat++) { // 仅支持2D数据教学 const values = samples.map(i => data[i][feat]).sort((a,b) => a-b); for (let i = 0; i < values.length - 1; i++) { const threshold = (values[i] + values[i+1]) / 2; const left = samples.filter(i => data[i][feat] < threshold); const right = samples.filter(i => data[i][feat] >= threshold); if (left.length === 0 || right.length === 0) continue; const gain = giniGain(labels, left, right); if (gain > bestGain) { bestGain = gain; bestSplit = { feat, threshold, left, right }; } } } if (!bestSplit) { return { type: 'leaf', label: majorityVote(labels.filter((_,i) => samples.includes(i))), count: samples.length }; } return { type: 'internal', feature: bestSplit.feat, threshold: bestSplit.threshold.toFixed(2), left: splitNode(bestSplit.left, depth + 1), right: splitNode(bestSplit.right, depth + 1), count: samples.length }; } return splitNode(data.map((_,i) => i), 0); } // D3渲染树(简化版) function renderTree(root, svg) { const treeLayout = d3.tree().size([400, 300]); const nodes = treeLayout(d3.hierarchy(root)); // 绑定节点点击事件:展开/折叠子树 svg.selectAll('.node') .data(nodes) .enter().append('g') .attr('class', 'node') .on('click', function(event, d) { if (d.children) { d._children = d.children; d.children = null; } else if (d._children) { d.children = d._children; d._children = null; } renderTree(root, svg); // 重新渲染 }); }关键设计:
buildDecisionTree()返回纯JS对象,而非TF.js tensor。因为树结构是离散逻辑,用JS递归比TF.js张量操作更直观易懂;且D3.tree只接受JS对象。分裂阈值保留两位小数(.toFixed(2)),避免显示3.2000000000000004这种反教学的数字。
4. 避坑指南:那些让教学演示当场翻车的12个细节
这个小程序看似只是“前端画图+JS计算”,但实际落地时,有12个高频坑能让演示在课堂上卡死、错乱或误导学生。以下按现象、原因、解决三步给出可立即复用的方案,全部来自真实课堂翻车记录。
4.1 现象:拖动学习率滑块,直线不动,console里报Cannot read property 'dataSync' of undefined
原因:regressor.trainStep()返回loss值,但若训练数据为空(如CSV上传失败后未重置xData/yData),trainStep()内部pred.sub(y)会返回null tensor,后续.dataSync()调用失败。
解决:在trainStep()开头加防御性检查:
trainStep(x, y) { if (x == null || y == null) { console.warn('Training data is null, skipping step'); return 0; } // 原逻辑... }4.2 现象:上传新CSV后,决策树节点文字重叠,无法阅读
原因:D3渲染树时,节点文本使用固定字体大小(如12px),但新数据范围变化导致树布局宽度压缩,文字挤在一起。
解决:动态计算字体大小,与树宽度成反比:
const fontSize = Math.max(8, Math.min(14, 400 / (maxDepth * 2))); // 宽度400px时,深度3用12px nodeEnter.append('text').style('font-size', `${fontSize}px`);4.3 现象:KNN模块中,当k设为1时,决策热力图大片空白
原因:热力图网格分辨率固定(如50×50),但k=1时,每个网格点只依赖最近1个样本,若该样本恰好远离网格点,预测结果不稳定,D3插值产生空白。
解决:k=1时改用最近邻插值(nearest neighbor),而非双线性插值:
const colorScale = d3.scaleOrdinal() .domain(['0','1']) .range(['#ff6b6b','#4ecdc4']); // k=1时用pointRadius=1,k>1时用radius=3 heatmap.attr('shape-rendering', k === 1 ? 'crispEdges' : 'geometricPrecision');4.4 现象:切换算法(如从线性回归切到KNN)后,旧图表残留,新图表叠加其上
原因:D3的selectAll().data().enter()模式未清理旧元素,新渲染的SVG元素追加到旧元素之后,造成视觉污染。
解决:每次切换算法前,清除对应SVG组的所有子元素:
d3.select('#plot-area').selectAll('*').remove(); // 清除全部 // 或更精准: d3.select('#regression-group').selectAll('*').remove(); d3.select('#knn-group').selectAll('*').remove();4.5 现象:移动端触摸滑块时,直线抖动严重,学生无法精确调节
原因:移动端input事件触发频率远高于桌面端,导致trainStep()被高频调用,参数剧烈震荡。
解决:添加防抖(debounce):
function debounce(func, wait) { let timeout; return function executedFunction() { const later = () => { clearTimeout(timeout); func(...arguments); }; clearTimeout(timeout); timeout = setTimeout(later, wait); }; } d3.select('#lr-slider').on('input', debounce(() => { regressor.learningRate = +this.value; regressor.trainStep(xData, yData); updatePlot(); }, 100)); // 100ms防抖5. 教学增强技巧:用“对比实验”设计让学生自己发现算法本质
光会调参不够,教学价值在于引导学生提出问题。我在西电机器学习期末复习课上,用三个对比实验让学生自己总结出算法特性,效果远超直接讲定义。
5.1 实验一:线性回归 vs 多项式回归——用同一组数据,看“过拟合”如何肉眼发生
准备一组带明显二次趋势的数据(如y = x^2 + noise),让学生先用线性回归拟合,再切换到二次多项式回归(y = w1*x + w2*x^2 + b)。关键操作:
- 步骤1:固定学习率=0.01,训练100步,观察线性回归直线始终无法贴合曲线;
- 步骤2:将学习率调至0.1,线性回归开始震荡,但多项式回归快速收敛;
- 步骤3:把数据中最后3个点改为异常值(outlier),线性回归直线大幅偏移,多项式回归波动更剧烈。
学生收获:不用讲“过拟合定义”,他们亲眼看到——增加模型复杂度(多项式)提升拟合能力,但也放大噪声敏感性;学习率不是越大越好,需匹配模型复杂度。
5.2 实验二:KNN的k值选择——用鸢尾花数据子集,画出“k-准确率曲线”
内置鸢尾花数据(150样本,3类),让学生:
- 上传数据后,用80%做训练,20%做测试;
- 拖动k滑块从1到15,每调一次,D3实时在右侧画布绘制当前k对应的测试准确率点;
- 观察曲线:k=1时准确率高但波动大(受噪声影响),k=5时达到峰值,k>10后准确率缓慢下降(欠拟合)。
表格:典型k值下的表现对比
| k值 | 决策边界特点 | 对异常值鲁棒性 | 计算开销 | 教学启示 |
|-----|--------------|----------------|----------|----------|
| 1 | 锯齿状,紧贴训练点 | 极差 | 低 | “最近邻”本质是记忆,非泛化 |
| 5 | 平滑,保留主要趋势 | 中等 | 中 | “投票”带来稳定性,需平衡k |
| 15 | 过于平滑,忽略局部结构 | 强 | 高 | k过大=用全局均值代替局部规律 |
5.3 实验三:决策树剪枝——用“深度vs叶节点数”散点图揭示奥卡姆剃刀
让学生调整最大深度(1~6),记录每次生成的叶节点数量和测试准确率:
- 深度=1:2个叶节点,准确率65%;
- 深度=3:8个叶节点,准确率82%;
- 深度=6:32个叶节点,准确率78%(过拟合)。
D3用气泡图展示:横轴深度,纵轴准确率,气泡大小=叶节点数。学生立刻看出——准确率在深度=3~4时达到平台,继续加深只增加节点数,不提升性能。这就是剪枝的直观依据。
我带过的每届学生,做完这三个实验后,期末考“解释过拟合原因”题的得分率从52%升到89%。不是因为他们记住了定义,而是他们在滑块拖动、数据上传、图表跳变的过程中,亲手触摸到了算法的呼吸节奏。这种肌肉记忆,比背十遍公式管用得多。希望帮到你。
本文还有配套的精品资源,点击获取