Learning rate

학습률

Set the step size of each gradient update.

···
html
<div class="stage"><div class="heading"><strong id="title"></strong><span id="status"></span></div><canvas id="plot"></canvas><div class="controls"><label id="control-label"></label><input id="parameter" type="range" min="0" max="100" value="55"><button id="restart" type="button">다시 보기</button></div></div>
css
.stage{position:absolute;inset:0;padding:10px 12px 8px;display:flex;flex-direction:column;gap:5px;background:var(--bg)}
.heading{display:flex;justify-content:space-between;align-items:center;gap:8px;font-size:12px;white-space:nowrap}.heading strong{overflow:hidden;text-overflow:ellipsis}.heading span{color:var(--muted);font-size:11px}
canvas{width:100%;min-height:0;flex:1;border:1px solid var(--line);border-radius:8px;background:var(--surface)}
.controls{display:flex;align-items:center;gap:10px;font-size:11px;color:var(--muted);min-height:28px}.controls label{min-width:80px}.controls input{flex:1;accent-color:var(--accent)}.controls button{border:1px solid var(--line);border-radius:6px;padding:4px 8px;background:var(--surface);color:var(--fg);cursor:pointer}
@media(max-width:500px){.stage{padding:7px 8px 6px}.controls{display:none}.heading{font-size:11px}.heading span{font-size:10px}}
js
const kind = "learning-rate", label = "학습률";
const canvas = document.getElementById('plot'), ctx = canvas.getContext('2d');
const title = document.getElementById('title'), status = document.getElementById('status');
const slider = document.getElementById('parameter');
title.textContent = label;
const parameterNames={
 'feature-engineering':'원본 값','normalization-scaling':'선택 값','train-test-split':'학습 비율',
 'cross-validation':'평가 폴드','overfitting':'모델 복잡도','bias-variance':'복잡도',
 'regularization':'벌점 강도','linear-regression':'기울기','logistic-regression':'경계 기울기',
 'gradient-descent':'진행 단계','learning-rate':'학습률','loss-function':'잔차 보기',
 'decision-tree':'분할 기준','random-forest':'입력 값','k-nearest-neighbors':'이웃 수',
 'confusion-matrix':'예측 결과','precision-recall':'임계값','roc-auc':'임계값',
 'neural-network':'전달 단계','attention':'단어 가중치','softmax-temperature':'온도',
 'nearest-neighbor-search':'탐색 단계'};
document.getElementById('control-label').textContent=parameterNames[kind];
const css = getComputedStyle(document.documentElement);
const color = (key) => css.getPropertyValue(key).trim();
const C = {fg:color('--fg'), muted:color('--muted'), line:color('--line'), a:color('--accent'), b:color('--accent-2'), c:color('--accent-3'), surface:color('--surface')};
let w=1,h=1, userValue=null, started=performance.now();
let seed=34729; function rand(){ seed=(seed+0x6D2B79F5)|0; let z=seed; z=Math.imul(z^(z>>>15),z|1); z^=z+Math.imul(z^(z>>>7),z|61); return ((z^(z>>>14))>>>0)/4294967296; }
const points=Array.from({length:30},(_,i)=>{ const x=.08+.84*rand(); return {x,y:.18+.63*x+(rand()-.5)*.26, cls:x+(rand()-.5)*.65>.52?1:0}; });
const groups=Array.from({length:26},(_,i)=>{ const k=i%3, a=rand()*6.28, r=Math.sqrt(rand())*.16; return {x:[.25,.72,.48][k]+Math.cos(a)*r,y:[.3,.31,.72][k]+Math.sin(a)*r,cls:k}; });
const X=(x)=>24+x*(w-48), Y=(y)=>h-22-y*(h-44);
function resize(){ const rect=canvas.getBoundingClientRect(); w=Math.max(1,rect.width);h=Math.max(1,rect.height); const d=Math.min(devicePixelRatio||1,2); canvas.width=Math.round(w*d);canvas.height=Math.round(h*d);ctx.setTransform(d,0,0,d,0,0); }
new ResizeObserver(resize).observe(canvas); resize();
function text(s,x,y,fill=C.muted,size=11,align='left'){ctx.fillStyle=fill;ctx.font='600 '+size+'px sans-serif';ctx.textAlign=align;ctx.fillText(s,x,y);ctx.textAlign='left';}
function line(x1,y1,x2,y2,stroke=C.line,width=1){ctx.strokeStyle=stroke;ctx.lineWidth=width;ctx.beginPath();ctx.moveTo(x1,y1);ctx.lineTo(x2,y2);ctx.stroke();}
function dot(x,y,fill=C.a,r=4){ctx.fillStyle=fill;ctx.beginPath();ctx.arc(X(x),Y(y),r,0,7);ctx.fill();}
function bar(x,y,width,height,fill){ctx.fillStyle=fill;ctx.fillRect(x,y,width,height);}
function path(fn,fill=C.a,width=2,start=0,end=1){ctx.strokeStyle=fill;ctx.lineWidth=width;ctx.beginPath();for(let i=0;i<=90;i++){let x=start+(end-start)*i/90,y=fn(x);if(i===0)ctx.moveTo(X(x),Y(y));else ctx.lineTo(X(x),Y(y));}ctx.stroke();}
function axes(){line(X(0),Y(0),X(1),Y(0));line(X(0),Y(0),X(0),Y(1));}
function matrix(values,labels){let size=Math.min(70,(w-70)/2,(h-55)/2),ox=w/2-size,oy=h/2-size;for(let r=0;r<2;r++)for(let c=0;c<2;c++){bar(ox+c*size+2,oy+r*size+2,size-4,size-4,(r===c?C.c:C.b));text(String(values[r*2+c]),ox+(c+.5)*size,oy+(r+.58)*size,C.surface,Math.max(13,Math.min(22,size*.32)),'center');}text(labels[0],ox+size/2,oy-5,C.muted,10,'center');text(labels[1],ox+size*1.5,oy-5,C.muted,10,'center');text('실제 +',ox-5,oy+size*.58,C.muted,10,'right');text('실제 −',ox-5,oy+size*1.58,C.muted,10,'right');}
function polyfit(data,degree){let n=degree+1,A=Array.from({length:n},()=>Array(n+1).fill(0));for(let r=0;r<n;r++){for(let c=0;c<n;c++)A[r][c]=data.reduce((s,p)=>s+Math.pow(p.x,r+c),0);A[r][n]=data.reduce((s,p)=>s+p.y*Math.pow(p.x,r),0);}for(let i=0;i<n;i++){let best=i;for(let r=i+1;r<n;r++)if(Math.abs(A[r][i])>Math.abs(A[best][i]))best=r;[A[i],A[best]]=[A[best],A[i]];let q=A[i][i]||1e-8;for(let c=i;c<=n;c++)A[i][c]/=q;for(let r=0;r<n;r++)if(r!==i){let f=A[r][i];for(let c=i;c<=n;c++)A[r][c]-=f*A[i][c];}}return A.map(row=>row[n]);}
function evalPoly(coefs,x){return coefs.reduce((s,a,i)=>s+a*Math.pow(x,i),0);}
const fitData=[{x:.08,y:.24},{x:.20,y:.42},{x:.32,y:.38},{x:.44,y:.59},{x:.56,y:.55},{x:.68,y:.72},{x:.80,y:.68},{x:.92,y:.83}];
function draw(now){
 const phase=((now-started)%5200)/5200, v=userValue===null ? .16+.74*phase : userValue/100;
 ctx.clearRect(0,0,w,h); let caption='';
 switch(kind){
 case 'feature-engineering': {axes();points.slice(0,16).forEach(p=>dot(p.x,p.y,C.muted,3));let px=points[Math.floor(v*15)];line(X(px.x),Y(0),X(px.x),Y(px.y),C.b,2);dot(px.x,px.y,C.a,6);text('원본: 길이 '+Math.round(px.x*100),X(.04),Y(.92));text('새 특성: 길이² '+Math.round(px.x*px.x*100),X(.04),Y(.79),C.a);caption='원본 값에서 새 열을 계산';break;}
 case 'normalization-scaling': {let vals=[12,35,78,23,61],selected=Math.min(4,Math.floor(v*5));vals.forEach((a,i)=>{let yy=Y(.85-i*.17),norm=(a-12)/66;text(String(a),X(.02),yy+4);bar(X(.18),yy-6,(w-48)*.6*a/80,12,C.muted);bar(X(.18),yy+8,(w-48)*.6*norm,5,i===selected?C.c:C.a);});caption=vals[selected]+' → '+((vals[selected]-12)/66).toFixed(2)+' · MinMax';break;}
 case 'train-test-split': {axes();points.forEach((p,i)=>dot(p.x,p.y,i<Math.round(v*points.length)?C.a:C.b,4));text('학습 '+Math.round(v*30),X(.03),Y(.93),C.a);text('평가 '+(30-Math.round(v*30)),X(.03),Y(.80),C.b);caption='평가 점은 학습에서 제외';break;}
 case 'cross-validation': {let fold=Math.min(4,Math.floor(v*5));for(let r=0;r<5;r++){let yy=Y(.82-r*.16);text(String(r+1),X(.03),yy+4);for(let c=0;c<5;c++)bar(X(.12+c*.16),yy-9,(w-48)*.145,18,c===r?C.b:C.a);if(r===fold){ctx.strokeStyle=C.fg;ctx.lineWidth=2;ctx.strokeRect(X(.11),yy-12,(w-48)*.81,24);}}text('평가 폴드 '+(fold+1),X(.12),Y(.05),C.b);caption='각 행에서 평가 폴드를 교체';break;}
 case 'overfitting': {axes();fitData.forEach(p=>dot(p.x,p.y,C.fg,4));let degree=v>.36?7:1,co=polyfit(fitData,degree);if(degree===7){let simple=polyfit(fitData,1);path(x=>evalPoly(simple,x),C.muted,1,.08,.92);}path(x=>evalPoly(co,x),degree===7?C.b:C.a,3,.08,.92);text(degree===7?'복잡한 곡선':'단순한 직선',X(.05),Y(.9),degree===7?C.b:C.a);caption='훈련 점을 과하게 따르면 흔들림';break;}
 case 'bias-variance': {axes();path(x=>.75-.52*x+.34*x*x,C.a,2);path(x=>.18+.65*x*x,C.b,2);let x=v;line(X(x),Y(0),X(x),Y(1),C.muted);dot(x,.75-.52*x+.34*x*x,C.a,5);dot(x,.18+.65*x*x,C.b,5);text('편향',X(.05),Y(.84),C.a);text('분산',X(.75),Y(.84),C.b);caption='복잡도가 커질수록 균형 이동';break;}
 case 'regularization': {axes();fitData.forEach(p=>dot(p.x,p.y,C.muted,3));let avg=fitData.reduce((s,p)=>s+p.y,0)/fitData.length, slope=.7*(1-v);path(x=>avg+slope*(x-.5),C.a,3);path(x=>avg+.7*(x-.5),C.line,2);text('계수 '+slope.toFixed(2),X(.06),Y(.9),C.a);caption='벌점이 커지면 계수가 작아짐';break;}
 case 'linear-regression': {axes();points.forEach(p=>dot(p.x,p.y,C.muted,3));let slope=.15+.55*v,intercept=.22;path(x=>intercept+slope*x,C.a,3);let mse=points.reduce((s,p)=>s+(p.y-intercept-slope*p.x)**2,0)/points.length;text('평균제곱오차 '+mse.toFixed(3),X(.04),Y(.91),C.a);caption='직선이 오차를 줄이며 이동';break;}
 case 'logistic-regression': {axes();points.forEach(p=>dot(p.x,p.cls? .78:.16,p.cls?C.a:C.b,4));let steep=3+v*13;path(x=>1/(1+Math.exp(-steep*(x-.52))),C.c,3);line(X(.52),Y(0),X(.52),Y(1),C.line);caption='시그모이드로 확률을 추정';break;}
 case 'gradient-descent': {axes();path(x=>.12+2.8*(x-.55)**2,C.a,3);let x=.95-.38*v;dot(x,.12+2.8*(x-.55)**2,C.b,7);line(X(x),Y(.12+2.8*(x-.55)**2),X(x-.08),Y(.12+2.8*(x-.55)**2),C.b,2);caption='기울기 반대 방향으로 한 걸음';break;}
 case 'learning-rate': {axes();path(x=>.12+2.8*(x-.55)**2,C.line,2);let lr=.02+v*.5,x=.95;for(let i=0;i<6;i++){dot(x,.12+2.8*(x-.55)**2,i===5?C.b:C.a,i===5?6:3);let next=x-lr*5.6*(x-.55);line(X(x),Y(.12+2.8*(x-.55)**2),X(Math.max(-1,Math.min(2,next))),Y(.12+2.8*(next-.55)**2),C.a);x=next;}caption='학습률 '+lr.toFixed(2)+' · 이동 폭';break;}
 case 'loss-function': {axes();let slope=.25+.7*v,mse=0;points.slice(0,12).forEach(p=>{let estimate=.18+slope*p.x;mse+=(p.y-estimate)**2;dot(p.x,p.y,C.muted,3);line(X(p.x),Y(p.y),X(p.x),Y(estimate),C.b,2);});path(x=>.18+slope*x,C.a,2);caption='평균제곱오차 '+(mse/12).toFixed(3);break;}
 case 'decision-tree': {axes();let split=.3+.4*v;line(X(split),Y(0),X(split),Y(1),C.a,3);line(X(split),Y(.51),X(1),Y(.51),C.b,3);groups.forEach(p=>dot(p.x,p.y,[C.a,C.b,C.c][p.cls],4));text('x < '+split.toFixed(2),X(.03),Y(.92),C.a);caption='질문으로 영역을 차례로 나눔';break;}
 case 'random-forest': {let votes=[v<.4?0:1,v<.65?0:1,v<.8?1:0],bw=(w-70)/3;votes.forEach((a,i)=>{let xx=28+i*bw;bar(xx,h*.23,bw-12,h*.42,a?C.a:C.b);text('나무 '+(i+1),xx+4,h*.2,C.muted,10);text(a?'A':'B',xx+(bw-12)/2,h*.5,a?'#fff':'#15151a',18,'center');});caption='세 나무 투표 → '+(votes.filter(Boolean).length>=2?'A':'B');break;}
 case 'k-nearest-neighbors': {axes();let q={x:.48,y:.50},sorted=groups.map(p=>({...p,d:Math.hypot(p.x-q.x,p.y-q.y)})).sort((a,b)=>a.d-b.d),k=Math.max(1,Math.round(1+v*8)),radius=sorted[k-1].d;ctx.strokeStyle=C.line;ctx.beginPath();ctx.ellipse(X(q.x),Y(q.y),radius*(w-48),radius*(h-44),0,0,7);ctx.stroke();groups.forEach(p=>dot(p.x,p.y,[C.a,C.b,C.c][p.cls],4));sorted.slice(0,k).forEach(p=>line(X(q.x),Y(q.y),X(p.x),Y(p.y),C.muted));dot(q.x,q.y,C.fg,7);caption='가까운 '+k+'개 이웃의 표결';break;}
 case 'confusion-matrix': {let n=Math.round(8*v);matrix([12+n,3,5,14-n],['예측 양성','예측 음성']);caption='행: 실제 · 열: 예측';break;}
 case 'precision-recall': {let threshold=v,selected=points.filter(p=>p.x>threshold),tp=selected.filter(p=>p.cls).length,actual=points.filter(p=>p.cls).length;axes();points.forEach(p=>dot(p.x,p.cls?.72:.26,p.x>threshold?(p.cls?C.c:C.b):C.muted,4));line(X(threshold),Y(0),X(threshold),Y(1),C.a,2);text('정밀도 '+(selected.length?Math.round(100*tp/selected.length):0)+'%',X(.04),Y(.92),C.c);text('재현율 '+Math.round(100*tp/actual)+'%',X(.55),Y(.92),C.a);caption='임계값 이동 → 두 지표 변화';break;}
 case 'roc-auc': {axes();line(X(0),Y(0),X(1),Y(1),C.line);let samples=points.map(p=>({score:p.x,cls:p.cls})).sort((a,b)=>b.score-a.score),pos=samples.filter(p=>p.cls).length,neg=samples.length-pos;let coords=[[0,0]],tp=0,fp=0;samples.forEach(p=>{if(p.cls)tp++;else fp++;coords.push([fp/neg,tp/pos]);});ctx.strokeStyle=C.a;ctx.lineWidth=3;ctx.beginPath();coords.forEach(([x,y],i)=>i?ctx.lineTo(X(x),Y(y)):ctx.moveTo(X(x),Y(y)));ctx.stroke();let auc=0;for(let i=1;i<coords.length;i++)auc+=(coords[i][0]-coords[i-1][0])*(coords[i][1]+coords[i-1][1])/2;let idx=Math.min(coords.length-1,Math.floor(v*(coords.length-1)));dot(coords[idx][0],coords[idx][1],C.b,7);text('AUC '+auc.toFixed(2),X(.58),Y(.12),C.a);caption='임계값별 TPR · FPR';break;}
 case 'neural-network': {let layers=[3,4,2],pulse=Math.min(1,Math.floor(v*3));layers.forEach((n,l)=>{for(let i=0;i<n;i++){let x=.18+l*.32,y=(i+1)/(n+1);if(l<2){let next=layers[l+1];for(let j=0;j<next;j++)line(X(x),Y(y),X(x+.32),Y((j+1)/(next+1)),l===pulse?C.a:C.line,l===pulse?2:1);}dot(x,y,l<=pulse?C.a:C.surface,7);ctx.strokeStyle=C.a;ctx.beginPath();ctx.arc(X(x),Y(y),7,0,7);ctx.stroke();}});caption='입력 → 은닉층 → 출력';break;}
 case 'attention': {let words=['이','글','의','주제'],weights=[.1,.18,.2+.5*v,.52-.5*v],sum=weights.reduce((a,b)=>a+b,0);words.forEach((word,i)=>{let xx=X(.08+i*.23),hh=(h-55)*weights[i]/sum;bar(xx,Y(0)-hh,Math.max(22,(w-48)*.16),hh,[C.line,C.a,C.b,C.c][i]);text(word,xx+8,Y(0)+14,C.fg,11);});caption='문맥에 따라 단어 가중치가 변함';break;}
 case 'softmax-temperature': {let logits=[2.2,1.4,.4],temp=.25+v*2.5,e=logits.map(x=>Math.exp((x-2.2)/temp)),sum=e.reduce((a,b)=>a+b,0);e.forEach((x,i)=>{let xx=X(.14+i*.28),hh=(h-58)*x/sum;bar(xx,Y(0)-hh,Math.max(25,(w-48)*.17),hh,[C.a,C.b,C.c][i]);text(Math.round(100*x/sum)+'%',xx,Y(0)-hh-5,C.fg,11);});caption='온도 '+temp.toFixed(2)+' · 확률 분포';break;}
 case 'nearest-neighbor-search': {axes();let q={x:.49,y:.49},visited=Math.max(1,Math.round(v*groups.length)),sorted=groups.map((p,i)=>({...p,i,d:Math.hypot(p.x-q.x,p.y-q.y)})),best=sorted.slice(0,visited).sort((a,b)=>a.d-b.d)[0];groups.forEach((p,i)=>dot(p.x,p.y,i<visited?C.muted:C.line,3));dot(best.x,best.y,C.b,7);dot(q.x,q.y,C.a,7);line(X(q.x),Y(q.y),X(best.x),Y(best.y),C.b,2);caption='후보 '+visited+'개 확인 · 현재 최근접';break;}
 }
 status.textContent=caption;
 requestAnimationFrame(draw);
}
slider.addEventListener('input',()=>{userValue=Number(slider.value);});
document.getElementById('restart').addEventListener('click',()=>{started=performance.now();userValue=null;slider.value='55';});
requestAnimationFrame(draw);

Learning rate is the stride length on a downhill walk. A short stride takes time; a long one can leap across the valley.

Each gradient update multiplies the gradient by the learning rate. It can stay fixed or follow a schedule that changes during training. The demo draws six steps on the same loss curve at different stride lengths.

A suitable value depends on model, batch size, and optimizer. Oscillating or exploding loss can indicate an excessive rate, though data and scaling issues deserve a check too.

When to use

Tune when learning is slow or unstable; compare loss curves under matched conditions.

Open as page ↗