import { TeachableTransformer } from './index.js'; import { TokenizerViz, AttentionMatrixViz, PoolingVectorViz, EmbeddingSpace2DViz } from './visualizer.js'; // Global classifier instance & Visualizer instances const classifier = new TeachableTransformer(); let tokenizerViz = null; let attentionViz = null; let poolingViz = null; let embeddingSpaceViz = null; let latestPredictionResult = null; function initVisualizers() { tokenizerViz = new TokenizerViz('tokenizer-viz-container'); attentionViz = new AttentionMatrixViz('attention-viz-container', 'attn-arcs-container'); poolingViz = new PoolingVectorViz('pooling-viz-container'); embeddingSpaceViz = new EmbeddingSpace2DViz('pca-canvas', 'pca-legend'); // Canvas click interaction for direct 2D vector space querying if (embeddingSpaceViz) { embeddingSpaceViz.setClickCallback((cx, cy) => { if (classifier.dataset.length === 0) { logToConsole('Dataset empty. Add training examples before querying vector space.', 'warn'); return; } logToConsole(`User queried 2D embedding space at canvas pos (${Math.round(cx)}, ${Math.round(cy)})`, 'info'); // Map 2D canvas click to nearest dataset points const norm = embeddingSpaceViz.fromCanvasCoords(cx, cy, embeddingSpaceViz.canvas.width, embeddingSpaceViz.canvas.height); const distances = classifier.dataset.map(item => ({ text: item.text, label: item.label, distance: Math.hypot(norm.x - item.x2d, norm.y - item.y2d) })); distances.sort((a, b) => a.distance - b.distance); const k = parseInt(document.getElementById('k-value')?.value, 10) || 3; const topK = distances.slice(0, Math.min(k, distances.length)); const votes = {}; topK.forEach(n => votes[n.label] = (votes[n.label] || 0) + 1); let bestLabel = null, maxV = -1; for (let [l, c] of Object.entries(votes)) { if (c > maxV) { maxV = c; bestLabel = l; } } const simResult = { predictedLabel: bestLabel, confidence: maxV / topK.length, nearestNeighbors: topK, inputEmbedding: topK[0] ? classifier.dataset.find(d => d.text === topK[0].text).embedding : new Float32Array(384), text: `2D Coordinates (${norm.x.toFixed(2)}, ${norm.y.toFixed(2)})` }; latestPredictionResult = simResult; displayPrediction(simResult.text, simResult); logToConsole(`Canvas Point Inference -> Predicted: '${bestLabel}' (${Math.round(simResult.confidence * 100)}% confidence)`, 'success'); }); } } // Helper to update visual log stream function logToConsole(message, type = 'info') { console.log(`[${type.toUpperCase()}] ${message}`); const logOutput = document.getElementById('log-output'); if (logOutput) { const entry = document.createElement('div'); entry.className = `log-entry log-${type}`; const time = new Date().toLocaleTimeString(); entry.innerHTML = `[${time}] ${escapeHtml(message)}`; logOutput.appendChild(entry); logOutput.scrollTop = logOutput.scrollHeight; } } function escapeHtml(text) { return text.replace(/&/g, "&").replace(//g, ">"); } function updateStatus(message, state = 'loading', progress = null) { const statusText = document.getElementById('status-text'); const statusDot = document.getElementById('status-dot'); const progressBar = document.getElementById('progress-bar'); const progressContainer = document.getElementById('progress-container'); if (statusText) statusText.textContent = message; if (statusDot) { statusDot.className = `status-dot ${state}`; } if (progressBar && progressContainer) { if (progress !== null && progress < 100) { progressContainer.style.display = 'block'; progressBar.style.width = `${progress}%`; } else { progressContainer.style.display = 'none'; } } } function updateModelSpecs(modelName, dim = null) { const specModelName = document.getElementById('spec-model-name'); const specEmbedDim = document.getElementById('spec-embed-dim'); if (specModelName) specModelName.textContent = modelName; if (specEmbedDim) { if (dim) { specEmbedDim.textContent = `${dim} dims`; } else { const knownDims = { 'Xenova/all-MiniLM-L6-v2': 384, 'Xenova/all-mpnet-base-v2': 768, 'Xenova/bge-small-en-v1.5': 384, 'Xenova/gte-small': 384, 'Xenova/distilbert-base-uncased': 768 }; specEmbedDim.textContent = `${knownDims[modelName] || '384'} dims`; } } } function renderDataset() { const container = document.getElementById('dataset-list'); const countBadge = document.getElementById('dataset-count'); if (!container) return; if (countBadge) countBadge.textContent = `${classifier.dataset.length} examples`; container.innerHTML = ''; if (classifier.dataset.length === 0) { container.innerHTML = '
No training examples added yet. Add examples above or click "Reset Demo Dataset".
'; if (embeddingSpaceViz) embeddingSpaceViz.updateData([], null); return; } const currentDim = classifier.dataset[0]?.embedding?.length || 384; classifier.dataset.forEach((item, index) => { const card = document.createElement('div'); card.className = 'example-card'; let labelClass = 'label-custom'; const labelLower = item.label.toLowerCase(); if (labelLower === 'positive') labelClass = 'label-positive'; else if (labelLower === 'negative') labelClass = 'label-negative'; else if (labelLower === 'neutral') labelClass = 'label-neutral'; card.innerHTML = `
${escapeHtml(item.label)}

"${escapeHtml(item.text)}"

${currentDim}-dimensional embedding cached
`; container.appendChild(card); }); // Add delete listeners container.querySelectorAll('.delete-btn').forEach(btn => { btn.addEventListener('click', (e) => { const idx = parseInt(e.target.getAttribute('data-index'), 10); const removed = classifier.dataset.splice(idx, 1)[0]; logToConsole(`Removed example: "${removed.text}"`, 'warn'); renderDataset(); }); }); // Update 2D PCA Embedding Canvas if (embeddingSpaceViz) { embeddingSpaceViz.updateData(classifier.dataset, latestPredictionResult); } } function displayPrediction(text, result) { latestPredictionResult = { text, ...result }; const resultBox = document.getElementById('result-box'); const labelEl = document.getElementById('predicted-label'); const confEl = document.getElementById('confidence-val'); const confBar = document.getElementById('confidence-fill'); const neighborsList = document.getElementById('neighbors-list'); const vectorPreview = document.getElementById('vector-preview'); if (!resultBox) return; resultBox.style.display = 'block'; // Label badge styling let labelClass = 'label-custom'; const labelLower = (result.predictedLabel || '').toLowerCase(); if (labelLower === 'positive') labelClass = 'label-positive'; else if (labelLower === 'negative') labelClass = 'label-negative'; else if (labelLower === 'neutral') labelClass = 'label-neutral'; labelEl.innerHTML = `${escapeHtml(result.predictedLabel)}`; const pct = Math.round(result.confidence * 100); confEl.textContent = `${pct}% (${result.nearestNeighbors.length} neighbor vote${result.nearestNeighbors.length > 1 ? 's' : ''})`; confBar.style.width = `${pct}%`; // Nearest neighbors neighborsList.innerHTML = ''; result.nearestNeighbors.forEach((neighbor, i) => { const item = document.createElement('div'); item.className = 'neighbor-item'; let nLabelClass = 'label-custom'; const nLabelLower = (neighbor.label || '').toLowerCase(); if (nLabelLower === 'positive') nLabelClass = 'label-positive'; else if (nLabelLower === 'negative') nLabelClass = 'label-negative'; else if (nLabelLower === 'neutral') nLabelClass = 'label-neutral'; const dist = neighbor.distance.toFixed(4); item.innerHTML = `
#${i + 1}
${escapeHtml(neighbor.label)} Euclidean Distance: ${dist}
"${escapeHtml(neighbor.text)}"
`; neighborsList.appendChild(item); }); // Vector visualization preview if (vectorPreview && result.inputEmbedding) { const dim = result.inputEmbedding.length; const previewDims = result.inputEmbedding.slice(0, 8).map(n => (n >= 0 ? '+' : '') + n.toFixed(3)).join(', '); vectorPreview.textContent = `Vector embedding [${dim} dims]: [${previewDims}, ...]`; updateModelSpecs(classifier.modelName, dim); } // Trigger Visual Explorers Updates if (tokenizerViz) tokenizerViz.render(text); if (attentionViz) attentionViz.render(text); if (poolingViz) poolingViz.render(result.inputEmbedding); if (embeddingSpaceViz) embeddingSpaceViz.updateData(classifier.dataset, latestPredictionResult); } export async function switchModel(modelName) { logToConsole(`Switching active model to '${modelName}'...`, 'info'); updateStatus(`Loading Transformer Model (${modelName})...`, 'loading', 10); classifier.modelName = modelName; updateModelSpecs(modelName); try { const startTime = performance.now(); await classifier.load((progress) => { if (progress.status === 'progress' && progress.total) { const pct = Math.round((progress.loaded / progress.total) * 100); updateStatus(`Downloading ${progress.file || 'model'} (${pct}%)...`, 'loading', pct); } }); const loadTime = ((performance.now() - startTime) / 1000).toFixed(2); logToConsole(`Model '${modelName}' loaded in ${loadTime}s!`, 'success'); updateStatus('Model loaded & ready', 'ready'); // Re-extract embeddings for existing training examples using new model if (classifier.dataset.length > 0) { logToConsole(`Re-computing embeddings for ${classifier.dataset.length} training examples with new model...`, 'info'); updateStatus('Re-extracting dataset embeddings...', 'loading'); for (let item of classifier.dataset) { item.embedding = await classifier.getEmbedding(item.text); } renderDataset(); logToConsole('Dataset re-embedded successfully.', 'success'); updateStatus('Model ready', 'ready'); } // Re-run prediction if test text exists const testInput = document.getElementById('test-text'); const predictBtn = document.getElementById('predict-btn'); if (testInput && testInput.value.trim() && predictBtn) { predictBtn.click(); } } catch (err) { logToConsole(`Error switching to model ${modelName}: ${err.message}`, 'error'); updateStatus(`Error: ${err.message}`, 'error'); } } export async function runDemo() { logToConsole('Starting Teachable Transformer demonstration...', 'info'); const modelSelect = document.getElementById('model-select'); const selectedModel = modelSelect ? modelSelect.value : 'Xenova/all-MiniLM-L6-v2'; classifier.modelName = selectedModel; updateModelSpecs(selectedModel); updateStatus(`Loading Model (${selectedModel})...`, 'loading', 10); try { // 1. Download model into browser cache const startTime = performance.now(); await classifier.load((progress) => { if (progress.status === 'progress' && progress.total) { const pct = Math.round((progress.loaded / progress.total) * 100); updateStatus(`Downloading ${progress.file || 'model'} (${pct}%)...`, 'loading', pct); } }); const loadTime = ((performance.now() - startTime) / 1000).toFixed(2); logToConsole(`Model loaded successfully in ${loadTime}s!`, 'success'); updateStatus('Model loaded & ready', 'ready'); // 2. "Train" the model (Transfer Learning via KNN) logToConsole('Populating initial sentiment dataset...', 'info'); classifier.dataset = []; // Reset dataset for clean demo run await classifier.addExample("I love this product, it's amazing!", "positive"); await classifier.addExample("This is the best day ever.", "positive"); await classifier.addExample("I hate this, it broke immediately.", "negative"); await classifier.addExample("Terrible experience, would not recommend.", "negative"); await classifier.addExample("Where is the nearest post office?", "neutral"); await classifier.addExample("I need to buy some milk later.", "neutral"); logToConsole(`Added ${classifier.dataset.length} training examples.`, 'success'); renderDataset(); // 3. Test the model on unseen data const testText = "I am so happy with my purchase!"; const k = 3; logToConsole(`Running inference for test text: "${testText}" (k=${k})...`, 'info'); const result = await classifier.predict(testText, k); console.log(`Text: "${testText}"`); console.log(`Predicted Label: ${result.predictedLabel}`); logToConsole(`Prediction: '${result.predictedLabel}' (Confidence: ${Math.round(result.confidence * 100)}%)`, 'success'); // Update UI test input field const testInput = document.getElementById('test-text'); if (testInput) testInput.value = testText; displayPrediction(testText, result); } catch (err) { logToConsole(`Error running demo: ${err.message}`, 'error'); updateStatus(`Error: ${err.message}`, 'error'); } } // Bind UI controls function setupUI() { initVisualizers(); const modelSelect = document.getElementById('model-select'); const addForm = document.getElementById('add-example-form'); const predictBtn = document.getElementById('predict-btn'); const resetBtn = document.getElementById('reset-demo-btn'); const clearBtn = document.getElementById('clear-dataset-btn'); const presetBtns = document.querySelectorAll('.preset-badge'); // Stepper click handlers const stepperSteps = document.querySelectorAll('.stepper-step'); stepperSteps.forEach(step => { step.addEventListener('click', () => { const stepNum = step.getAttribute('data-step'); stepperSteps.forEach(s => s.classList.remove('active')); step.classList.add('active'); const targetCard = document.getElementById(`viz-step-${stepNum}`); if (targetCard) { targetCard.scrollIntoView({ behavior: 'smooth', block: 'center' }); targetCard.classList.add('pulse-highlight'); setTimeout(() => targetCard.classList.remove('pulse-highlight'), 1200); } }); }); // Animate Forward Pass Button const animBtn = document.getElementById('animate-pipeline-btn'); if (animBtn) { animBtn.addEventListener('click', () => { animBtn.disabled = true; logToConsole('Animating inference forward pass through pipeline stages...', 'info'); const steps = [1, 2, 3, 4]; let i = 0; const interval = setInterval(() => { if (i < steps.length) { const stepNum = steps[i]; stepperSteps.forEach(s => s.classList.remove('active')); const curStep = document.querySelector(`.stepper-step[data-step="${stepNum}"]`); if (curStep) curStep.classList.add('active'); const targetCard = document.getElementById(`viz-step-${stepNum}`); if (targetCard) { targetCard.scrollIntoView({ behavior: 'smooth', block: 'center' }); targetCard.classList.add('pulse-highlight'); setTimeout(() => targetCard.classList.remove('pulse-highlight'), 1000); } i++; } else { clearInterval(interval); animBtn.disabled = false; logToConsole('Forward pass walkthrough completed.', 'success'); } }, 1200); }); } if (modelSelect) { modelSelect.addEventListener('change', async (e) => { const newModel = e.target.value; await switchModel(newModel); }); } if (presetBtns) { presetBtns.forEach(btn => { btn.addEventListener('click', () => { const labelInput = document.getElementById('example-label'); if (labelInput) { labelInput.value = btn.getAttribute('data-label'); } }); }); } if (addForm) { addForm.addEventListener('submit', async (e) => { e.preventDefault(); const textInput = document.getElementById('example-text'); const labelInput = document.getElementById('example-label'); const text = textInput.value.trim(); const label = labelInput.value.trim(); if (!text || !label) return; const submitBtn = addForm.querySelector('button[type="submit"]'); submitBtn.disabled = true; submitBtn.textContent = 'Extracting...'; try { updateStatus(`Extracting features for "${text.slice(0, 25)}..."`, 'loading'); await classifier.addExample(text, label); logToConsole(`Added example for label '${label}': "${text}"`, 'success'); renderDataset(); textInput.value = ''; updateStatus('Model ready', 'ready'); } catch (err) { logToConsole(`Error adding example: ${err.message}`, 'error'); } finally { submitBtn.disabled = false; submitBtn.textContent = '+ Add Example'; } }); } if (predictBtn) { predictBtn.addEventListener('click', async () => { const testInput = document.getElementById('test-text'); const kInput = document.getElementById('k-value'); const text = testInput.value.trim(); const k = parseInt(kInput.value, 10) || 3; if (!text) return; predictBtn.disabled = true; predictBtn.textContent = 'Classifying...'; updateStatus('Embedding text & searching nearest neighbors...', 'loading'); try { const startTime = performance.now(); const result = await classifier.predict(text, k); const elapsed = (performance.now() - startTime).toFixed(1); logToConsole(`Inference completed in ${elapsed}ms -> '${result.predictedLabel}'`, 'success'); displayPrediction(text, result); updateStatus('Model ready', 'ready'); } catch (err) { logToConsole(`Prediction error: ${err.message}`, 'error'); alert(err.message); } finally { predictBtn.disabled = false; predictBtn.textContent = 'Classify Text'; } }); } if (resetBtn) { resetBtn.addEventListener('click', async () => { logToConsole('Resetting to default demo dataset...', 'info'); await runDemo(); }); } if (clearBtn) { clearBtn.addEventListener('click', () => { classifier.dataset = []; logToConsole('Dataset cleared.', 'warn'); renderDataset(); }); } } // Run setup and demo when page loads if (typeof document !== 'undefined') { if (document.readyState === 'loading') { window.addEventListener('DOMContentLoaded', () => { setupUI(); runDemo(); }); } else { setupUI(); runDemo(); } }