ujjwalpardeshi/pytorch-training-debugger
2
1<!DOCTYPE html>2<html lang="en">3<head>4<meta charset="UTF-8">5<meta name="viewport" content="width=device-width, initial-scale=1.0">6<title>PyTorch Training Debugger — Live Dashboard</title>7<script src="https://cdn.plot.ly/plotly-2.27.0.min.js"></script>8<style>9* { margin: 0; padding: 0; box-sizing: border-box; }10body { font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif; background: #0d1117; color: #c9d1d9; }11.header { background: #161b22; padding: 16px 24px; border-bottom: 1px solid #30363d; display: flex; align-items: center; gap: 16px; }12.header h1 { font-size: 20px; font-weight: 600; }13.header .status { padding: 4px 12px; border-radius: 12px; font-size: 13px; font-weight: 500; }14.status.connected { background: #238636; color: #fff; }15.status.disconnected { background: #da3633; color: #fff; }16.grid { display: grid; grid-template-columns: 1fr 1fr; grid-template-rows: 1fr 1fr; gap: 12px; padding: 12px; height: calc(100vh - 60px); }17.panel { background: #161b22; border: 1px solid #30363d; border-radius: 8px; overflow: hidden; display: flex; flex-direction: column; }18.panel-title { padding: 10px 16px; font-size: 14px; font-weight: 600; color: #58a6ff; border-bottom: 1px solid #30363d; background: #0d1117; }19.panel-body { flex: 1; padding: 8px; position: relative; min-height: 0; }20.panel-body > div:first-child { width: 100%; height: 100%; }21.placeholder { display: flex; align-items: center; justify-content: center; height: 100%; color: #484f58; font-style: italic; }22#controls { display: flex; gap: 8px; align-items: center; }23#controls select, #controls button { background: #21262d; color: #c9d1d9; border: 1px solid #30363d; padding: 6px 12px; border-radius: 6px; cursor: pointer; font-size: 13px; }24#controls button:hover { background: #30363d; }25#controls button.primary { background: #238636; border-color: #238636; color: #fff; }26#summary { padding: 16px; font-size: 13px; line-height: 1.8; overflow-y: auto; }27#summary .row { display: flex; justify-content: space-between; border-bottom: 1px solid #21262d; padding: 4px 0; }28#summary .label { color: #8b949e; }29#summary .value { font-weight: 600; }30#summary .score { font-size: 24px; color: #58a6ff; text-align: center; margin: 12px 0; }31.actions-list { display: flex; flex-wrap: wrap; gap: 4px; margin-top: 8px; }32.action-tag { padding: 2px 8px; border-radius: 4px; font-size: 11px; font-weight: 500; }33.action-tag.investigate { background: #1f6feb33; color: #58a6ff; }34.action-tag.fix { background: #23863633; color: #3fb950; }35.action-tag.terminal { background: #da363333; color: #f85149; }36.action-tag.wrong { background: #da363366; color: #f85149; }37</style>38</head>39<body>40<div class="header">41 <h1>PyTorch Training Debugger</h1>42 <div id="connStatus" class="status disconnected">Disconnected</div>43 <div id="controls">44 <select id="taskSelect">45 <option value="task_001">Task 1 — Exploding Gradients (Easy)</option>46 <option value="task_002">Task 2 — Vanishing Gradients (Easy)</option>47 <option value="task_003">Task 3 — Data Leakage (Medium)</option>48 <option value="task_004">Task 4 — Overfitting (Medium)</option>49 <option value="task_005">Task 5 — BatchNorm Eval (Hard)</option>50 <option value="task_006">Task 6 — Code Bug (Hard)</option>51 <option value="task_007">Task 7 — Scheduler Misconfigured (Med-Hard)</option>52 </select>53 <button class="primary" onclick="runBaseline()">Run Baseline</button>54 </div>55</div>56<div class="grid">57 <div class="panel">58 <div class="panel-title">Training Metrics</div>59 <div class="panel-body"><div id="metricsChart"><div class="placeholder">Run baseline to see metrics</div></div></div>60 </div>61 <div class="panel">62 <div class="panel-title">Gradient & Weight Heatmap</div>63 <div class="panel-body"><div id="gradientChart"><div class="placeholder">Not yet inspected</div></div></div>64 </div>65 <div class="panel">66 <div class="panel-title">Action Timeline & Rewards</div>67 <div class="panel-body"><div id="timelineChart"><div class="placeholder">No actions yet</div></div></div>68 </div>69 <div class="panel">70 <div class="panel-title">Episode Summary</div>71 <div class="panel-body" id="summary">72 <div class="placeholder">Waiting for episode</div>73 </div>74 </div>75</div>76 77<script>78const host = window.location.host;79const wsProto = window.location.protocol === 'https:' ? 'wss:' : 'ws:';80let ws = null;81let actions = [];82let rewards = [];83let cumRewards = [];84let obs = null;85 86function setStatus(connected) {87 const el = document.getElementById('connStatus');88 el.textContent = connected ? 'Connected' : 'Disconnected';89 el.className = 'status ' + (connected ? 'connected' : 'disconnected');90}91 92function connect() {93 ws = new WebSocket(`${wsProto}//${host}/ws`);94 ws.onopen = () => setStatus(true);95 ws.onclose = () => { setStatus(false); setTimeout(connect, 2000); };96 ws.onerror = () => ws.close();97 ws.onmessage = (ev) => {98 const msg = JSON.parse(ev.data);99 if (msg.type === 'observation' && msg.data) {100 // Framework wraps: {type: "observation", data: {observation: {...}, reward, done}}101 const wrapper = msg.data;102 const obsData = wrapper.observation || wrapper;103 obsData.reward = wrapper.reward;104 obsData.done = wrapper.done;105 handleObservation(obsData);106 }107 };108}109 110function handleObservation(data) {111 obs = data;112 if (data.reward !== null && data.reward !== undefined) {113 rewards.push(data.reward);114 const prev = cumRewards.length > 0 ? cumRewards[cumRewards.length - 1] : 0;115 cumRewards.push(prev + data.reward);116 }117 if (data.episode_state && data.episode_state.actions_taken) {118 actions = data.episode_state.actions_taken;119 }120 updateMetrics(data);121 updateGradients(data);122 updateTimeline();123 updateSummary(data);124}125 126function updateMetrics(d) {127 const traces = [];128 if (d.training_loss_history && d.training_loss_history.length > 0) {129 const valid = d.training_loss_history.filter(v => isFinite(v));130 traces.push({ y: valid, name: 'Train Loss', line: { color: '#f85149' } });131 }132 if (d.val_loss_history && d.val_loss_history.length > 0) {133 const valid = d.val_loss_history.filter(v => isFinite(v));134 traces.push({ y: valid, name: 'Val Loss', line: { color: '#f0883e', dash: 'dash' } });135 }136 if (d.val_accuracy_history && d.val_accuracy_history.length > 0) {137 traces.push({ y: d.val_accuracy_history, name: 'Val Accuracy', yaxis: 'y2', line: { color: '#3fb950' } });138 }139 if (traces.length === 0) return;140 Plotly.newPlot('metricsChart', traces, {141 paper_bgcolor: 'transparent', plot_bgcolor: 'transparent',142 font: { color: '#c9d1d9', size: 11 },143 margin: { t: 10, b: 30, l: 50, r: 50 },144 xaxis: { title: 'Epoch', gridcolor: '#21262d' },145 yaxis: { title: 'Loss', gridcolor: '#21262d' },146 yaxis2: { title: 'Accuracy', overlaying: 'y', side: 'right', range: [0, 1], gridcolor: '#21262d' },147 legend: { x: 0, y: 1.15, orientation: 'h' },148 showlegend: true,149 }, { responsive: true });150}151 152function updateGradients(d) {153 if (!d.gradient_stats || d.gradient_stats.length === 0) return;154 const layers = d.gradient_stats.map(g => g.layer_name);155 const norms = d.gradient_stats.map(g => g.mean_norm);156 const colors = d.gradient_stats.map(g => g.is_exploding ? '#f85149' : g.is_vanishing ? '#1f6feb' : '#3fb950');157 Plotly.newPlot('gradientChart', [{158 x: layers, y: norms, type: 'bar',159 marker: { color: colors },160 text: d.gradient_stats.map(g => g.is_exploding ? 'EXPLODING' : g.is_vanishing ? 'VANISHING' : 'Normal'),161 textposition: 'auto',162 }], {163 paper_bgcolor: 'transparent', plot_bgcolor: 'transparent',164 font: { color: '#c9d1d9', size: 11 },165 margin: { t: 10, b: 30, l: 50, r: 20 },166 yaxis: { title: 'Mean Grad Norm', gridcolor: '#21262d', type: 'log' },167 xaxis: { gridcolor: '#21262d' },168 }, { responsive: true });169}170 171function updateTimeline() {172 if (actions.length === 0) return;173 const colors = actions.map(a => {174 if (a.startsWith('inspect')) return '#1f6feb';175 if (a.startsWith('fix') || a === 'modify_config' || a === 'patch_data_loader' || a === 'add_callback' || a === 'replace_optimizer') return '#238636';176 if (a.startsWith('mark_diagnosed')) return '#da3633';177 if (a === 'restart_run') return '#f0883e';178 return '#484f58';179 });180 Plotly.newPlot('timelineChart', [181 { x: actions.map((_, i) => i + 1), y: rewards, type: 'bar', name: 'Step Reward', marker: { color: rewards.map(r => r >= 0 ? '#3fb950' : '#f85149') } },182 { x: actions.map((_, i) => i + 1), y: cumRewards, type: 'scatter', name: 'Cumulative', line: { color: '#58a6ff', width: 2 } }183 ], {184 paper_bgcolor: 'transparent', plot_bgcolor: 'transparent',185 font: { color: '#c9d1d9', size: 11 },186 margin: { t: 10, b: 30, l: 50, r: 20 },187 xaxis: { title: 'Step', gridcolor: '#21262d', tickvals: actions.map((_, i) => i + 1), ticktext: actions.map(a => a.split(':')[0].replace('inspect_', 'i_').replace('mark_diagnosed', 'diag')) },188 yaxis: { title: 'Reward', gridcolor: '#21262d' },189 legend: { x: 0, y: 1.15, orientation: 'h' },190 }, { responsive: true });191}192 193function updateSummary(d) {194 const s = d.episode_state || {};195 const avail = d.available_actions || [];196 let html = '';197 if (d.done) {198 html += `<div class="score">Episode Complete</div>`;199 }200 html += '<div class="row"><span class="label">Task</span><span class="value">' + (d.run_id || '-') + '</span></div>';201 html += '<div class="row"><span class="label">Steps</span><span class="value">' + (s.step_count || 0) + '</span></div>';202 html += '<div class="row"><span class="label">Gradients Inspected</span><span class="value">' + (s.gradients_inspected ? 'Yes' : 'No') + '</span></div>';203 html += '<div class="row"><span class="label">Gradients Normal</span><span class="value">' + (s.gradients_were_normal ? 'Yes' : '-') + '</span></div>';204 html += '<div class="row"><span class="label">Data Inspected</span><span class="value">' + (s.data_inspected ? 'Yes' : 'No') + '</span></div>';205 html += '<div class="row"><span class="label">Model Modes Inspected</span><span class="value">' + (s.model_modes_inspected ? 'Yes' : 'No') + '</span></div>';206 html += '<div class="row"><span class="label">Code Inspected</span><span class="value">' + (s.code_inspected ? 'Yes' : 'No') + '</span></div>';207 html += '<div class="row"><span class="label">Fix Applied</span><span class="value">' + (s.fix_action_taken ? 'Yes' : 'No') + '</span></div>';208 html += '<div class="row"><span class="label">Restarted</span><span class="value">' + (s.restart_after_fix ? 'Yes' : 'No') + '</span></div>';209 html += '<div class="row"><span class="label">Diagnosed</span><span class="value">' + (s.diagnosis_submitted ? 'Yes' : 'No') + '</span></div>';210 if (d.code_snippet) {211 html += '<div style="margin-top:12px"><span class="label">Code:</span><pre style="background:#0d1117;padding:8px;border-radius:4px;font-size:11px;overflow:auto;max-height:120px;margin-top:4px">' + d.code_snippet.code.replace(/</g,'<') + '</pre></div>';212 }213 html += '<div style="margin-top:8px"><span class="label">Available Actions:</span></div>';214 html += '<div class="actions-list">';215 avail.forEach(a => {216 let cls = 'investigate';217 if (a.startsWith('fix') || a === 'modify_config' || a === 'patch_data_loader' || a === 'add_callback' || a === 'replace_optimizer') cls = 'fix';218 if (a === 'mark_diagnosed' || a === 'restart_run') cls = 'terminal';219 html += `<span class="action-tag ${cls}">${a}</span>`;220 });221 html += '</div>';222 document.getElementById('summary').innerHTML = html;223}224 225function sendStep(action) {226 return new Promise(resolve => {227 const handler = (ev) => {228 const msg = JSON.parse(ev.data);229 if (msg.type === 'observation') {230 ws.removeEventListener('message', handler);231 resolve(msg);232 }233 };234 ws.addEventListener('message', handler);235 ws.send(JSON.stringify({ type: 'step', data: action }));236 });237}238 239function sendReset(taskId) {240 return new Promise(resolve => {241 const handler = (ev) => {242 const msg = JSON.parse(ev.data);243 if (msg.type === 'observation') {244 ws.removeEventListener('message', handler);245 resolve(msg);246 }247 };248 ws.addEventListener('message', handler);249 ws.send(JSON.stringify({ type: 'reset', data: { task_id: taskId, seed: 42 } }));250 });251}252 253async function runBaseline() {254 const taskId = document.getElementById('taskSelect').value;255 actions = []; rewards = []; cumRewards = [];256 if (!ws || ws.readyState !== WebSocket.OPEN) return;257 258 const delay = (ms) => new Promise(r => setTimeout(r, ms));259 260 // Reset261 await sendReset(taskId);262 await delay(300);263 264 // Step 1: Inspect gradients265 await sendStep({ action_type: 'inspect_gradients' });266 await delay(300);267 268 const gs = obs && obs.gradient_stats ? obs.gradient_stats : [];269 const anyExploding = gs.some(g => g.is_exploding);270 const anyVanishing = gs.some(g => g.is_vanishing);271 272 if (anyExploding) {273 await sendStep({ action_type: 'modify_config', target: 'learning_rate', value: 0.001 });274 await delay(300);275 await sendStep({ action_type: 'restart_run' });276 await delay(300);277 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'lr_too_high' });278 return;279 }280 281 if (anyVanishing) {282 await sendStep({ action_type: 'modify_config', target: 'learning_rate', value: 0.01 });283 await delay(300);284 await sendStep({ action_type: 'restart_run' });285 await delay(300);286 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'vanishing_gradients' });287 return;288 }289 290 // Step 2: Inspect data291 await sendStep({ action_type: 'inspect_data_batch' });292 await delay(300);293 294 const dbs = obs && obs.data_batch_stats ? obs.data_batch_stats : {};295 if (dbs.class_overlap_score && dbs.class_overlap_score > 0.5) {296 await sendStep({ action_type: 'patch_data_loader' });297 await delay(300);298 await sendStep({ action_type: 'restart_run' });299 await delay(300);300 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'data_leakage' });301 return;302 }303 304 // Check for overfitting (train loss low, val loss rising)305 const tl = obs && obs.training_loss_history ? obs.training_loss_history : [];306 const vl = obs && obs.val_loss_history ? obs.val_loss_history : [];307 const lastTrainLoss = tl.length > 0 ? tl[tl.length - 1] : 999;308 const lastValLoss = vl.length > 0 ? vl[vl.length - 1] : 0;309 const earlyValLoss = vl.length > 5 ? vl[5] : lastValLoss;310 const isOverfitting = lastTrainLoss < 0.1 && lastValLoss > earlyValLoss;311 312 if (isOverfitting) {313 await sendStep({ action_type: 'modify_config', target: 'weight_decay', value: 0.01 });314 await delay(300);315 await sendStep({ action_type: 'restart_run' });316 await delay(300);317 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'overfitting' });318 return;319 }320 321 // Step 3: Inspect model modes322 await sendStep({ action_type: 'inspect_model_modes' });323 await delay(300);324 325 const modes = obs && obs.model_mode_info ? obs.model_mode_info : {};326 const anyEval = Object.values(modes).some(m => m === 'eval');327 if (anyEval) {328 await sendStep({ action_type: 'fix_model_mode' });329 await delay(300);330 await sendStep({ action_type: 'restart_run' });331 await delay(300);332 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'batchnorm_eval_mode' });333 return;334 }335 336 // Step 4: Inspect code337 await sendStep({ action_type: 'inspect_code' });338 await delay(300);339 340 if (obs && obs.code_snippet && obs.code_snippet.code) {341 const code = obs.code_snippet.code;342 const lines = code.split('\n');343 let fixLine = null, fixReplacement = null;344 for (let i = 0; i < lines.length; i++) {345 const ln = lines[i].trim();346 if (ln.includes('model.eval()')) { fixLine = i + 1; fixReplacement = lines[i].replace('model.eval()', 'model.train()'); break; }347 if (ln.includes('.detach()') && ln.includes('criterion')) { fixLine = i + 1; fixReplacement = lines[i].replace('.detach()', ''); break; }348 if (ln.includes('inplace=True')) { fixLine = i + 1; fixReplacement = lines[i].replace('inplace=True', ''); break; }349 }350 if (fixLine) {351 await sendStep({ action_type: 'fix_code', line: fixLine, replacement: fixReplacement });352 await delay(300);353 } else {354 // zero_grad_missing — find optimizer.step() and add zero_grad before it355 for (let i = 0; i < lines.length; i++) {356 if (lines[i].trim().includes('optimizer.step()')) {357 fixLine = i + 1;358 fixReplacement = ' optimizer.zero_grad()\n' + lines[i];359 break;360 }361 }362 if (fixLine) {363 await sendStep({ action_type: 'fix_code', line: fixLine, replacement: fixReplacement });364 await delay(300);365 }366 }367 await sendStep({ action_type: 'restart_run' });368 await delay(300);369 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'code_bug' });370 return;371 }372 373 // Step 5: Check for scheduler issue374 const va = obs && obs.val_accuracy_history ? obs.val_accuracy_history : [];375 const midAcc = va.length > 10 ? va[9] : 0;376 const endAcc = va.length > 0 ? va[va.length - 1] : 0;377 const stagnated = midAcc > 0.3 && (endAcc - midAcc) < 0.05;378 379 if (stagnated) {380 await sendStep({ action_type: 'modify_config', target: 'learning_rate', value: 0.005 });381 await delay(300);382 await sendStep({ action_type: 'restart_run' });383 await delay(300);384 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'scheduler_misconfigured' });385 return;386 }387 388 // Fallback389 await sendStep({ action_type: 'modify_config', target: 'weight_decay', value: 0.01 });390 await delay(300);391 await sendStep({ action_type: 'restart_run' });392 await delay(300);393 await sendStep({ action_type: 'mark_diagnosed', diagnosis: 'overfitting' });394}395 396connect();397</script>398</body>399</html>400 