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.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();
}
}