Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
dashboard.html400 linesDownload Raw Back to server
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,'&lt;') + '</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