const PYODIDE_URL = "https://cdn.jsdelivr.net/pyodide/v0.26.1/full/pyodide.mjs"; const RANGE = { min: -5, max: 5 }; const COLORS = { A: "#2d6cdf", B: "#e86f32" }; const REGION = { A: "rgba(45,108,223,.12)", B: "rgba(232,111,50,.12)" }; const BUILTIN_DATASETS = { "easy-linear": [ { x: -3.2, y: -2.0, label: "A" }, { x: -2.2, y: 0.1, label: "A" }, { x: -3.0, y: 2.3, label: "A" }, { x: 3.1, y: -2.1, label: "B" }, { x: 2.1, y: 0.2, label: "B" }, { x: 3.2, y: 2.2, label: "B" }, ], "hard-linear": [ { x: -3.4, y: -1.08, label: "A" }, { x: -0.2, y: 0.80, label: "A" }, { x: 3.2, y: 2.95, label: "A" }, { x: -3.2, y: -2.30, label: "B" }, { x: 0.2, y: -0.16, label: "B" }, { x: 3.4, y: 1.77, label: "B" }, ], nonlinear: [ { x: -0.7, y: -0.5, label: "A" }, { x: 0.7, y: -0.5, label: "A" }, { x: 0.0, y: 0.8, label: "A" }, { x: -3.2, y: -2.4, label: "B" }, { x: 3.2, y: -2.4, label: "B" }, { x: 0.0, y: 3.3, label: "B" }, ], }; const state = { points: structuredClone(BUILTIN_DATASETS['easy-linear']), datasetMode: "easy-linear", redoPoints: [], hoveredPoint: -1, label: "A", model: null, pendingUpdate: null, metrics: null, step: 0, permutation: [], cursor: 0, nextBatch: [], rngState: 42, history: [], viewIndex: -1, previewVersion: 0, previewPromise: null, busy: false, animating: false, ready: false }; const canvas = document.querySelector("#plot"); const ctx = canvas.getContext("2d"); const metricsCanvas = document.querySelector("#metrics-chart"); const metricsCtx = metricsCanvas.getContext("2d"); const ui = Object.fromEntries([...document.querySelectorAll("[id]")].map(el => [el.id, el])); function modelOptions() { return { loss: ui.loss.value, learning_rate: Number(ui["learning-rate"].value), batch_size: Math.max(1, Number.parseInt(ui["batch-size"].value, 10) || 1), batch_selection: ui["batch-selection"].value, random_state: randomSeed(), }; } function randomSeed() { return (Number.parseInt(ui["random-seed"].value, 10) || 0) >>> 0; } function resetModel(message = "Model reset. Take a step to begin training.") { state.model = null; state.pendingUpdate = null; state.metrics = null; state.step = 0; state.permutation = []; state.cursor = 0; state.nextBatch = []; state.history = []; state.viewIndex = -1; state.rngState = randomSeed(); queueNextBatch(); ui["step-count"].textContent = "Step 0"; renderView(); setMessage(message); updateControls(); refreshPreview(); } function setMessage(text, kind = "") { ui.message.textContent = text; ui.message.className = `message ${kind}`; } function hasBothClasses() { return state.points.some(p => p.label === "A") && state.points.some(p => p.label === "B"); } function updateControls() { ui.undo.disabled = state.points.length === 0 || state.busy || state.animating; ui.redo.disabled = state.redoPoints.length === 0 || state.busy || state.animating; ui.clear.disabled = state.points.length === 0 || state.busy || state.animating; ui["label-a"].disabled = state.busy || state.animating; ui["label-b"].disabled = state.busy || state.animating; for (const id of ["dataset-select", "loss", "learning-rate", "batch-size", "random-seed", "animation-speed", "batch-selection"]) ui[id].disabled = state.busy || state.animating; ui["history-prev"].disabled = state.viewIndex <= 0 || state.busy || state.animating; ui["history-reset"].disabled = state.viewIndex <= 0 || state.busy || state.animating; ui["history-next"].disabled = !state.ready || !hasBothClasses() || state.busy || state.animating; ui.animate.disabled = !state.ready || !hasBothClasses() || state.busy || state.animating; ui.pause.disabled = !state.animating; ui["point-count"].textContent = `${state.points.length} point${state.points.length === 1 ? "" : "s"}`; ui["empty-hint"].hidden = state.points.length > 0; //ui["plot-instructions"].textContent = state.datasetMode === "custom" ? "Click inside the plot to add a point with the selected label." : "Hover over a point to inspect its coordinates."; if (!hasBothClasses()) setMessage("Add at least one point from each class to begin."); } function shuffleIndices(count) { const values = Array.from({ length: count }, (_, i) => i); for (let i = values.length - 1; i > 0; i--) { const j = Math.floor(seededRandom() * (i + 1)); [values[i], values[j]] = [values[j], values[i]]; } return values; } function seededRandom() { state.rngState = (Math.imul(1664525, state.rngState) + 1013904223) >>> 0; return state.rngState / 4294967296; } function nextBatch() { const { batch_size, batch_selection } = modelOptions(); const size = Math.min(batch_size, state.points.length); if (batch_selection === "random") return shuffleIndices(state.points.length).slice(0, size); if (!state.permutation.length || state.cursor >= state.permutation.length) { state.permutation = shuffleIndices(state.points.length); state.cursor = 0; } const batch = state.permutation.slice(state.cursor, state.cursor + size); state.cursor += size; return batch; } function queueNextBatch() { state.nextBatch = hasBothClasses() ? nextBatch() : []; const count = state.nextBatch.length; ui["batch-summary"].textContent = count ? `Next batch: ${count} point${count === 1 ? "" : "s"}` : "No batch selected"; } function formatNumber(value) { if (Math.abs(value) < 0.0000005) return "0.0000"; if (Math.abs(value) >= 1000 || Math.abs(value) < 0.001) return value.toExponential(3); return value.toFixed(4); } function formatVector(vector) { return `[${vector.map(formatNumber).join(", ")}]`; } function visibleSnapshot() { return state.history[state.viewIndex] || { model: state.model, pendingUpdate: state.pendingUpdate, metrics: state.metrics, nextBatch: state.nextBatch, step: state.step }; } function renderParameters(snapshot = visibleSnapshot()) { ui["model-parameters"].textContent = snapshot.model ? `w_A = ${formatVector(snapshot.model.coef[0])}\nb_A = ${formatNumber(snapshot.model.intercept[0])}` : "w_A = —\nb_A = —"; ui["model-update"].textContent = snapshot.pendingUpdate ? `Δw_A = ${formatVector(snapshot.pendingUpdate.coef[0])}\nΔb_A = ${formatNumber(snapshot.pendingUpdate.intercept[0])}` : "Δw_A = —\nΔb_A = —"; } function captureSnapshot() { const snapshot = structuredClone({ model: state.model, pendingUpdate: state.pendingUpdate, metrics: state.metrics, nextBatch: state.nextBatch, step: state.step }); if (state.history.at(-1)?.step === state.step) state.history[state.history.length - 1] = snapshot; else state.history.push(snapshot); state.viewIndex = state.history.length - 1; renderView(); } function renderView() { const snapshot = visibleSnapshot(); renderParameters(snapshot); ui["step-count"].textContent = `Step ${snapshot.step}`; ui["history-position"].textContent = `Step ${snapshot.step} of ${Math.max(0, state.history.length - 1)}`; const count = snapshot.nextBatch.length; ui["batch-summary"].textContent = count ? `Next batch: ${count} point${count === 1 ? "" : "s"}` : "No batch selected"; updateControls(); draw(); drawMetrics(); } async function refreshPreview(reportError = true) { const version = ++state.previewVersion; if (!state.ready || !hasBothClasses() || !state.nextBatch.length) { renderView(); return; } state.previewPromise = (async () => { try { if (!state.model) { const pyOptions = state.pyodide.toPy(modelOptions()); state.pyodide.globals.set("model_options", pyOptions); try { state.pyodide.runPython("model = SGDClassifier(**model_options)"); } finally { pyOptions.destroy(); state.pyodide.globals.delete("model_options"); } } const X = state.nextBatch.map(i => [state.points[i].x, state.points[i].y]); const y = state.nextBatch.map(i => state.points[i].label); const allX = state.points.map(point => [point.x, point.y]); const allY = state.points.map(point => point.label); const pyX = state.pyodide.toPy(X), pyY = state.pyodide.toPy(y); const pyAllX = state.pyodide.toPy(allX), pyAllY = state.pyodide.toPy(allY); state.pyodide.globals.set("preview_X", pyX); state.pyodide.globals.set("preview_y", pyY); state.pyodide.globals.set("all_X", pyAllX); state.pyodide.globals.set("all_y", pyAllY); try { const result = state.pyodide.runPython(` preview_coef, preview_intercept = model.get_update(preview_X, preview_y, classes=["A", "B"]) metrics = model.evaluate(all_X, all_y) {"coef": model.coef_.tolist(), "intercept": model.intercept_.tolist(), "update_coef": preview_coef.tolist(), "update_intercept": preview_intercept.tolist(), "loss": metrics["loss"], "accuracy": metrics["accuracy"]} `); const values = result.toJs({ dict_converter: Object.fromEntries }); result.destroy(); if (version !== state.previewVersion) return; state.model = { coef: values.coef, intercept: values.intercept }; state.pendingUpdate = { coef: values.update_coef, intercept: values.update_intercept }; state.metrics = { loss: values.loss, accuracy: values.accuracy }; captureSnapshot(); } finally { pyX.destroy(); pyY.destroy(); pyAllX.destroy(); pyAllY.destroy(); for (const name of ["preview_X", "preview_y", "all_X", "all_y"]) state.pyodide.globals.delete(name); } } catch (error) { if (version === state.previewVersion) { state.pendingUpdate = null; renderView(); if (reportError) setMessage(error instanceof Error ? error.message : String(error), "error"); } } })(); return state.previewPromise; } async function initializeRuntime() { try { const { loadPyodide } = await import(PYODIDE_URL); const pyodide = await loadPyodide({ indexURL: "https://cdn.jsdelivr.net/pyodide/v0.26.1/full/" }); await pyodide.loadPackage("numpy"); const classifierUrl = new URL("sgd_classifier.py", window.location.href); classifierUrl.searchParams.set("v", Date.now().toString()); const source = await fetch(classifierUrl, { cache: "no-store" }).then(response => { if (!response.ok) throw new Error("Classifier source could not be loaded."); return response.text(); }); pyodide.FS.writeFile("sgd_classifier.py", source); pyodide.runPython("from sgd_classifier import SGDClassifier"); state.pyodide = pyodide; state.ready = true; ui["runtime-status"].className = "status ready"; ui["runtime-status"].innerHTML = "Python runtime ready"; updateControls(); refreshPreview(); } catch (error) { ui["runtime-status"].className = "status error"; ui["runtime-status"].innerHTML = "Runtime failed"; setMessage(error instanceof Error ? error.message : String(error), "error"); } } async function trainOneStep(fromAnimation = false) { if (!fromAnimation && ui["history-next"].disabled) return false; if (state.viewIndex < state.history.length - 1) { state.viewIndex += 1; renderView(); return true; } state.busy = true; updateControls(); ui["history-next"].querySelector(".button-text").textContent = "Training…"; try { if (state.previewPromise) await state.previewPromise; if (!state.model || !state.pendingUpdate) throw new Error("The parameter update is not ready yet."); const indices = state.nextBatch.slice(); const X = indices.map(i => [state.points[i].x, state.points[i].y]); const y = indices.map(i => state.points[i].label); const pyX = state.pyodide.toPy(X), pyY = state.pyodide.toPy(y); state.pyodide.globals.set("batch_X", pyX); state.pyodide.globals.set("batch_y", pyY); try { const result = state.pyodide.runPython(` model.train_step(batch_X, batch_y, classes=["A", "B"]) {"coef": model.coef_.tolist(), "intercept": model.intercept_.tolist()} `); state.model = result.toJs({ dict_converter: Object.fromEntries }); result.destroy(); } finally { pyX.destroy(); pyY.destroy(); state.pyodide.globals.delete("batch_X"); state.pyodide.globals.delete("batch_y"); } state.step += 1; queueNextBatch(); await refreshPreview(false); setMessage(`Completed SGD step ${state.step}. The next update is shown.`, "success"); renderView(); return true; } catch (error) { setMessage(error instanceof Error ? error.message : String(error), "error"); return false; } finally { state.busy = false; ui["history-next"].querySelector(".button-text").textContent = "Next"; updateControls(); } } async function startAnimation() { if (ui.animate.disabled) return; state.animating = true; updateControls(); setMessage("Animation running…", "success"); while (state.animating) { const succeeded = await trainOneStep(true); if (!succeeded) { state.animating = false; break; } await new Promise(resolve => setTimeout(resolve, Number(ui["animation-speed"].value))); } updateControls(); if (!state.animating && !ui.message.classList.contains("error")) setMessage("Animation paused."); } function pauseAnimation() { state.animating = false; updateControls(); } function drawMetrics() { const rect = metricsCanvas.getBoundingClientRect(); if (!rect.width || !rect.height) return; const ratio = window.devicePixelRatio || 1; metricsCanvas.width = Math.round(rect.width * ratio); metricsCanvas.height = Math.round(rect.height * ratio); metricsCtx.setTransform(ratio, 0, 0, ratio, 0, 0); metricsCtx.clearRect(0, 0, rect.width, rect.height); const data = state.history.filter(snapshot => snapshot.metrics); const pad = { left: 34, right: 34, top: 10, bottom: 22 }; const width = rect.width - pad.left - pad.right, height = rect.height - pad.top - pad.bottom; const bottom = pad.top + height; metricsCtx.font = "9px system-ui"; metricsCtx.fillStyle = "#77847d"; metricsCtx.strokeStyle = "#e4e9e5"; metricsCtx.lineWidth = 1; for (let fraction = 0; fraction <= 1; fraction += 0.5) { const y = pad.top + height * fraction; metricsCtx.beginPath(); metricsCtx.moveTo(pad.left,y); metricsCtx.lineTo(pad.left+width,y); metricsCtx.stroke(); } if (!data.length) { metricsCtx.textAlign="center"; metricsCtx.fillText("Metrics appear after initialization",rect.width/2,rect.height/2); return; } const maxLoss = Math.max(...data.map(snapshot => snapshot.metrics.loss), 1e-9) * 1.08; const pointX = index => data.length === 1 ? pad.left + width / 2 : pad.left + index / (data.length - 1) * width; const lossY = value => bottom - value / maxLoss * height; const accuracyY = value => bottom - value * height; metricsCtx.textAlign="right"; metricsCtx.fillStyle="#8b5cf6"; metricsCtx.fillText(maxLoss.toFixed(2),pad.left-5,pad.top+3); metricsCtx.fillText("0",pad.left-5,bottom+3); metricsCtx.textAlign="left"; metricsCtx.fillStyle="#159267"; metricsCtx.fillText("100%",pad.left+width+5,pad.top+3); metricsCtx.fillText("0%",pad.left+width+5,bottom+3); metricsCtx.textAlign="center"; metricsCtx.fillStyle="#77847d"; metricsCtx.fillText("step",pad.left+width/2,rect.height-3); const drawSeries = (color, getY) => { metricsCtx.strokeStyle=color; metricsCtx.lineWidth=2; metricsCtx.beginPath(); data.forEach((snapshot,index)=>{const x=pointX(index),y=getY(snapshot.metrics);if(index)metricsCtx.lineTo(x,y);else metricsCtx.moveTo(x,y);}); metricsCtx.stroke(); data.forEach((snapshot,index)=>{const x=pointX(index),y=getY(snapshot.metrics),active=index===state.viewIndex;metricsCtx.beginPath();metricsCtx.arc(x,y,active?5:2.5,0,Math.PI*2);metricsCtx.fillStyle=active?"white":color;metricsCtx.fill();metricsCtx.strokeStyle=color;metricsCtx.lineWidth=active?2.5:1;metricsCtx.stroke();}); }; drawSeries("#8b5cf6", metrics => lossY(metrics.loss)); drawSeries("#159267", metrics => accuracyY(metrics.accuracy)); } function dimensions() { const rect = canvas.getBoundingClientRect(), ratio = window.devicePixelRatio || 1; if (canvas.width !== Math.round(rect.width * ratio) || canvas.height !== Math.round(rect.height * ratio)) { canvas.width = Math.round(rect.width * ratio); canvas.height = Math.round(rect.height * ratio); } ctx.setTransform(ratio, 0, 0, ratio, 0, 0); return { width: rect.width, height: rect.height, pad: 40 }; } function toPixel(x, y, d) { const sizeX = d.width - d.pad * 2, sizeY = d.height - d.pad * 2; return [d.pad + (x - RANGE.min) / 10 * sizeX, d.height - d.pad - (y - RANGE.min) / 10 * sizeY]; } function fromPixel(px, py, d) { return [RANGE.min + (px - d.pad) / (d.width - d.pad * 2) * 10, RANGE.min + (d.height - d.pad - py) / (d.height - d.pad * 2) * 10]; } function predictedClass(x, y, model) { if (!model) return null; const scoreA = model.coef[0][0] * x + model.coef[0][1] * y + model.intercept[0]; const scoreB = model.coef[1][0] * x + model.coef[1][1] * y + model.intercept[1]; return scoreA >= scoreB ? "A" : "B"; } function draw() { const snapshot = visibleSnapshot(), model = snapshot.model, step = snapshot.step; const d = dimensions(); ctx.clearRect(0, 0, d.width, d.height); const [left, bottom] = toPixel(RANGE.min, RANGE.min, d), [right, top] = toPixel(RANGE.max, RANGE.max, d); ctx.fillStyle = "#fbfcfa"; ctx.fillRect(left, top, right - left, bottom - top); if (model && step > 0) { const cells = 55, cellW = (right-left)/cells, cellH = (bottom-top)/cells; for (let row=0; row1e-10) { const y1=-(w0*RANGE.min+b)/w1,y2=-(w0*RANGE.max+b)/w1,[px1,py1]=toPixel(RANGE.min,y1,d),[px2,py2]=toPixel(RANGE.max,y2,d); ctx.save();ctx.beginPath();ctx.rect(left,top,right-left,bottom-top);ctx.clip();ctx.strokeStyle="#26342c";ctx.lineWidth=2;ctx.beginPath();ctx.moveTo(px1,py1);ctx.lineTo(px2,py2);ctx.stroke();ctx.restore(); } else if(Math.abs(w0)>1e-10) { const x=-b/w0,[px]=toPixel(x,0,d);ctx.strokeStyle="#26342c";ctx.lineWidth=2;ctx.beginPath();ctx.moveTo(px,top);ctx.lineTo(px,bottom);ctx.stroke(); } } const selected = new Set(snapshot.nextBatch); state.points.forEach((point, index) => { const [x,y]=toPixel(point.x,point.y,d); if (selected.has(index)) { ctx.beginPath(); ctx.arc(x,y,11,0,Math.PI*2); ctx.fillStyle="rgba(255, 202, 58, .28)"; ctx.fill(); ctx.strokeStyle="#d29400"; ctx.lineWidth=2; ctx.stroke(); } ctx.beginPath();ctx.arc(x,y,6,0,Math.PI*2);ctx.fillStyle=COLORS[point.label];ctx.fill();ctx.strokeStyle="white";ctx.lineWidth=2;ctx.stroke(); }); if (state.hoveredPoint >= 0 && state.points[state.hoveredPoint]) { const point = state.points[state.hoveredPoint], [x,y] = toPixel(point.x,point.y,d); const text = `Class ${point.label} · (${point.x.toFixed(2)}, ${point.y.toFixed(2)})`; ctx.font = "11px system-ui"; const width = ctx.measureText(text).width + 16; const boxX = Math.min(Math.max(left + 4, x + 12), right - width - 4); const boxY = y - 35 < top ? y + 13 : y - 35; ctx.fillStyle = "rgba(31, 43, 36, .94)"; ctx.beginPath(); ctx.roundRect(boxX, boxY, width, 27, 5); ctx.fill(); ctx.fillStyle = "white"; ctx.fillText(text, boxX + 8, boxY + 18); } } canvas.addEventListener("mousemove", event => { const d=dimensions(), rect=canvas.getBoundingClientRect(), px=event.clientX-rect.left, py=event.clientY-rect.top; let nearest=-1, distance=10; state.points.forEach((point,index)=>{const [x,y]=toPixel(point.x,point.y,d), current=Math.hypot(px-x,py-y);if(current=0?"pointer":state.datasetMode==="custom"?"crosshair":"default"; draw(); } }); canvas.addEventListener("mouseleave",()=>{if(state.hoveredPoint>=0){state.hoveredPoint=-1;canvas.style.cursor=state.datasetMode==="custom"?"crosshair":"default";draw();}}); canvas.addEventListener("click", event => { if (state.busy || state.animating || state.datasetMode !== "custom") return; const d=dimensions(), rect=canvas.getBoundingClientRect(), px=event.clientX-rect.left,py=event.clientY-rect.top; if(pxd.width-d.pad||pyd.height-d.pad)return; const [x,y]=fromPixel(px,py,d);state.points.push({x,y,label:state.label});state.redoPoints=[];state.hoveredPoint=-1;resetModel("Dataset changed. The model was reset."); }); ui["label-a"].addEventListener("click",()=>{state.label="A";ui["label-a"].classList.add("active");ui["label-b"].classList.remove("active");}); ui["label-b"].addEventListener("click",()=>{state.label="B";ui["label-b"].classList.add("active");ui["label-a"].classList.remove("active");}); ui.undo.addEventListener("click",()=>{const point=state.points.pop();if(point)state.redoPoints.push(point);state.hoveredPoint=-1;resetModel("Dataset changed. The model was reset.");}); ui.redo.addEventListener("click",()=>{const point=state.redoPoints.pop();if(point)state.points.push(point);resetModel("Point restored. The model was reset.");}); ui.clear.addEventListener("click",()=>{state.points=[];state.redoPoints=[];state.hoveredPoint=-1;resetModel("Dataset cleared.");}); function loadDataset(mode, message) { ui["dataset-select"].value = mode; state.datasetMode = mode; ui["custom-data-controls"].hidden = state.datasetMode !== "custom"; state.points = state.datasetMode === "custom" ? [] : structuredClone(BUILTIN_DATASETS[state.datasetMode]); state.redoPoints = []; state.hoveredPoint = -1; canvas.style.cursor = state.datasetMode === "custom" ? "crosshair" : "default"; resetModel(message || (state.datasetMode === "custom" ? "Click the plot to create a custom dataset." : "Builtin dataset loaded. The model was reset.")); } ui["dataset-select"].addEventListener("change", () => loadDataset(ui["dataset-select"].value)); function selectTab(name) { for (const tab of ["configure", "train", "examples", "usage"]) { const active = tab === name; ui[`${tab}-tab`].classList.toggle("active", active); ui[`${tab}-panel`].classList.toggle("active", active); } if (name === "train") requestAnimationFrame(drawMetrics); } ui["configure-tab"].addEventListener("click", () => selectTab("configure")); ui["train-tab"].addEventListener("click", () => selectTab("train")); ui["examples-tab"].addEventListener("click", () => selectTab("examples")); ui["usage-tab"].addEventListener("click", () => selectTab("usage")); for (const button of document.querySelectorAll(".example-button")) { button.addEventListener("click", () => { pauseAnimation(); ui.loss.value = button.dataset.loss || "perceptron"; ui["learning-rate"].value = button.dataset.learningRate || "1"; ui["batch-size"].value = "1"; ui["batch-selection"].value = "random"; ui["random-seed"].value = "42"; ui["animation-speed"].value = "200"; ui["sampling-help"].textContent = "Each point is used once before the dataset is shuffled again."; loadDataset(button.dataset.dataset, `${button.textContent.trim()} example loaded.`); selectTab("train"); }); } ui["history-next"].addEventListener("click", trainOneStep); ui["history-prev"].addEventListener("click", () => { if (state.viewIndex > 0) { state.viewIndex -= 1; renderView(); } }); ui["history-reset"].addEventListener("click", () => { if (state.history.length) { state.viewIndex = 0; renderView(); } }); ui.animate.addEventListener("click", startAnimation); ui.pause.addEventListener("click", pauseAnimation); for (const id of ["loss","learning-rate","batch-size","random-seed","batch-selection"]) ui[id].addEventListener("change",()=>{if(id==="batch-selection")ui["sampling-help"].textContent=ui[id].value==="random"?"Each step independently samples a subset without replacement.":"Each point is used once before the dataset is shuffled again.";resetModel("Configuration changed. The model was reset.");}); window.addEventListener("keydown",event=>{if(event.code==="Space"&&!event.repeat&&!/INPUT|SELECT|TEXTAREA/.test(document.activeElement.tagName)){event.preventDefault();trainOneStep();}}); new ResizeObserver(draw).observe(canvas.parentElement); new ResizeObserver(drawMetrics).observe(metricsCanvas); queueNextBatch(); updateControls(); draw(); initializeRuntime();