CoolFace
Apppublic

lerobot/robot-learning-tutorial

sourceHugging Faceupdated 1y agoView on Hugging Face
508likes
line.py277 linesDownload Raw Back to plotly
1import plotly.graph_objects as go2import plotly.io as pio3import numpy as np4import os5import uuid6 7"""8Interactive line chart example (Baseline / Improved / Target) with a live slider.9 10Context: research-style training curves for multiple datasets (CIFAR-10, CIFAR-100, ImageNet-1K).11The slider "Augmentation α" blends the Improved curve between the Baseline (α=0)12and an augmented counterpart (α=1) via a simple mixing equation.13Export remains responsive, with no zoom and no mode bar.14"""15 16# Grid (x) and parameterization17N = 24018x = np.linspace(0, 1, N)19 20# Logistic helper for smooth learning curves21def logistic(xv: np.ndarray, ymin: float, ymax: float, k: float, x0: float) -> np.ndarray:22    return ymin + (ymax - ymin) / (1.0 + np.exp(-k * (xv - x0)))23 24# Plausible dataset params (baseline vs augmented) + a constant target line25datasets_params = [26    {27        "name": "CIFAR-10",28        "base": {"ymin": 0.10, "ymax": 0.90, "k": 10.0, "x0": 0.55},29        "aug":  {"ymin": 0.15, "ymax": 0.96, "k": 12.0, "x0": 0.40},30        "target": 0.97,31    },32    {33        "name": "CIFAR-100",34        "base": {"ymin": 0.05, "ymax": 0.70, "k": 9.5, "x0": 0.60},35        "aug":  {"ymin": 0.08, "ymax": 0.80, "k": 11.0, "x0": 0.45},36        "target": 0.85,37    },38    {39        "name": "ImageNet-1K",40        "base": {"ymin": 0.02, "ymax": 0.68, "k": 8.5, "x0": 0.65},41        "aug":  {"ymin": 0.04, "ymax": 0.75, "k": 9.5, "x0": 0.50},42        "target": 0.82,43    },44]45 46# Initial dataset index and alpha47alpha0 = 0.748ds0 = datasets_params[0]49base0 = logistic(x, **ds0["base"])50aug0 = logistic(x, **ds0["aug"])51target0 = np.full_like(x, ds0["target"], dtype=float)52 53# Traces: Baseline (fixed), Improved (blended by α), Target (constant goal)54blend = lambda l, e, a: (1 - a) * l + a * e55y1 = base056y2 = blend(base0, aug0, alpha0)57y3 = target058 59color_base = "#64748b"     # slate-50060color_improved = "#F981D4" # pink61color_target = "#4b5563"   # gray-600 (dash)62 63fig = go.Figure()64fig.add_trace(65    go.Scatter(66        x=x,67        y=y1,68        name="Baseline",69        mode="lines",70        line=dict(color=color_base, width=2, shape="spline", smoothing=0.6),71        hovertemplate="<b>%{fullData.name}</b><br>x=%{x:.2f}<br>y=%{y:.3f}<extra></extra>",72        showlegend=True,73    )74)75fig.add_trace(76    go.Scatter(77        x=x,78        y=y2,79        name="Improved",80        mode="lines",81        line=dict(color=color_improved, width=2, shape="spline", smoothing=0.6),82        hovertemplate="<b>%{fullData.name}</b><br>x=%{x:.2f}<br>y=%{y:.3f}<extra></extra>",83        showlegend=True,84    )85)86fig.add_trace(87    go.Scatter(88        x=x,89        y=y3,90        name="Target",91        mode="lines",92        line=dict(color=color_target, width=2, dash="dash"),93        hovertemplate="<b>%{fullData.name}</b><br>x=%{x:.2f}<br>y=%{y:.3f}<extra></extra>",94        showlegend=True,95    )96)97 98fig.update_layout(99    autosize=True,100    paper_bgcolor="rgba(0,0,0,0)",101    plot_bgcolor="rgba(0,0,0,0)",102    margin=dict(l=40, r=28, t=20, b=40),103    hovermode="x unified",104    legend=dict(105        orientation="v",106        x=1,107        y=0,108        xanchor="right",109        yanchor="bottom",110        bgcolor="rgba(255,255,255,0)",111        borderwidth=0,112    ),113    hoverlabel=dict(114        bgcolor="white",115        font=dict(color="#111827", size=12),116        bordercolor="rgba(0,0,0,0.15)",117        align="left",118        namelength=-1,119    ),120    xaxis=dict(121        showgrid=False,122        zeroline=False,123        showline=True,124        linecolor="rgba(0,0,0,0.25)",125        linewidth=1,126        ticks="outside",127        ticklen=6,128        tickcolor="rgba(0,0,0,0.25)",129        tickfont=dict(size=12, color="rgba(0,0,0,0.55)"),130        title=None,131        automargin=True,132        fixedrange=True,133    ),134    yaxis=dict(135        showgrid=False,136        zeroline=False,137        showline=True,138        linecolor="rgba(0,0,0,0.25)",139        linewidth=1,140        ticks="outside",141        ticklen=6,142        tickcolor="rgba(0,0,0,0.25)",143        tickfont=dict(size=12, color="rgba(0,0,0,0.55)"),144        title=None,145        tickformat=".2f",146        rangemode="tozero",147        automargin=True,148        fixedrange=True,149    ),150)151 152# Write the fragment next to this file into src/fragments/line.html (robust path)153output_path = os.path.join(os.path.dirname(__file__), "fragments", "line.html")154os.makedirs(os.path.dirname(output_path), exist_ok=True)155 156# Inject a small post-render script to round the hover box corners157post_script = """158(function(){159  function attach(gd){160    function round(){161      try {162        var root = gd && gd.parentNode ? gd.parentNode : document;163        var rects = root.querySelectorAll('.hoverlayer .hovertext rect');164        rects.forEach(function(r){ r.setAttribute('rx', 8); r.setAttribute('ry', 8); });165      } catch(e) {}166    }167    if (gd && gd.on) {168      gd.on('plotly_hover', round);169      gd.on('plotly_unhover', round);170      gd.on('plotly_relayout', round);171    }172    setTimeout(round, 0);173  }174  var plots = document.querySelectorAll('.js-plotly-plot');175  plots.forEach(attach);176})();177"""178 179html_plot = pio.to_html(180    fig,181    include_plotlyjs=False,182    full_html=False,183    post_script=post_script,184    config={185        "displayModeBar": False,186        "responsive": True,187        "scrollZoom": False,188        "doubleClick": False,189        "modeBarButtonsToRemove": [190            "zoom2d", "pan2d", "select2d", "lasso2d",191            "zoomIn2d", "zoomOut2d", "autoScale2d", "resetScale2d",192            "toggleSpikelines"193        ],194    },195)196 197# Build a self-contained fragment with a live slider (no mouseup required)198uid = uuid.uuid4().hex[:8]199slider_id = f"line-ex-alpha-{uid}"200container_id = f"line-ex-container-{uid}"201 202slider_tpl = '''203<div id="__CID__">204  __PLOT__205  <div class="plotly_controls" style="margin-top:12px; display:flex; gap:16px; align-items:center;">206    <label style="font-size:12px;color:rgba(0,0,0,.65); display:flex; align-items:center; gap:6px; white-space:nowrap; padding:6px 10px;">207      Dataset208      <select id="__DSID__" style="font-size:12px; padding:2px 6px;">209        <option value="0">CIFAR-10</option>210        <option value="1">CIFAR-100</option>211        <option value="2">ImageNet-1K</option>212      </select>213    </label>214    <label style="font-size:12px;color:rgba(0,0,0,.65);display:flex;align-items:center;gap:10px; flex:1; padding:6px 10px;">215      Augmentation α216      <input id="__SID__" type="range" min="0" max="1" step="0.01" value="__A0__" style="flex:1;">217      <span class="alpha-value">__A0__</span>218    </label>219  </div>220</div>221<script>222(function(){223  var container = document.getElementById('__CID__');224  if(!container) return;225  var gd = container.querySelector('.js-plotly-plot');226  var slider = document.getElementById('__SID__');227  var dsSelect = document.getElementById('__DSID__');228  var valueEl = container.querySelector('.alpha-value');229  var N = __N__;230  var xs = Array.from({length: N}, function(_,i){ return i/(N-1); });231  function logistic(x, ymin, ymax, k, x0){ return ymin + (ymax - ymin) / (1 + Math.exp(-k*(x - x0))); }232  function blend(l,e,a){ return (1-a)*l + a*e; }233  var datasets = [234    { name:'CIFAR-10',  base:{ymin:0.10,ymax:0.90,k:10.0,x0:0.55}, aug:{ymin:0.15,ymax:0.96,k:12.0,x0:0.40}, target:0.97 },235    { name:'CIFAR-100', base:{ymin:0.05,ymax:0.70,k:9.5,x0:0.60},  aug:{ymin:0.08,ymax:0.80,k:11.0,x0:0.45},  target:0.85 },236    { name:'ImageNet-1K', base:{ymin:0.02,ymax:0.68,k:8.5,x0:0.65}, aug:{ymin:0.04,ymax:0.75,k:9.5,x0:0.50},  target:0.82 }237  ];238  var dsi = 0;239  var yb = xs.map(function(x){ return logistic(x, datasets[dsi].base.ymin, datasets[dsi].base.ymax, datasets[dsi].base.k, datasets[dsi].base.x0); });240  var ya = xs.map(function(x){ return logistic(x, datasets[dsi].aug.ymin, datasets[dsi].aug.ymax, datasets[dsi].aug.k, datasets[dsi].aug.x0); });241  var yt = xs.map(function(){ return datasets[dsi].target; });242  function applyAlpha(a){243    var yi = yb.map(function(v,i){ return blend(v, ya[i], a); });244    Plotly.restyle(gd, {y:[yi]}, [1]); // only Improved changes with α245    if(valueEl) valueEl.textContent = a.toFixed(2);246  }247  function applyDataset(){248    var d = datasets[dsi];249    yb = xs.map(function(x){ return logistic(x, d.base.ymin, d.base.ymax, d.base.k, d.base.x0); });250    ya = xs.map(function(x){ return logistic(x, d.aug.ymin, d.aug.ymax, d.aug.k, d.aug.x0); });251    yt = xs.map(function(){ return d.target; });252    var a = parseFloat(slider.value)||0;253    var yi = yb.map(function(v,i){ return blend(v, ya[i], a); });254    Plotly.restyle(gd, {y:[yb]}, [0]); // Baseline255    Plotly.restyle(gd, {y:[yi]}, [1]); // Improved (blended)256    Plotly.restyle(gd, {y:[yt]}, [2]); // Target257  }258  var initA = parseFloat(slider.value)||0;259  slider.addEventListener('input', function(e){ applyAlpha(parseFloat(e.target.value)||0); });260  dsSelect.addEventListener('change', function(e){ dsi = parseInt(e.target.value)||0; applyDataset(); });261  setTimeout(function(){ applyDataset(); applyAlpha(initA); }, 0);262})();263</script>264'''265 266slider_html = (slider_tpl267    .replace('__CID__', container_id)268    .replace('__SID__', slider_id)269    .replace('__A0__', f"{alpha0:.2f}")270    .replace('__N__', str(N))271    .replace('__PLOT__', html_plot)272)273 274with open("../../app/src/content/fragments/line.html", "w", encoding="utf-8") as f:275    f.write(slider_html)276 277