경사 하강법

Gradient descent

손실이 낮아지는 방향으로 파라미터를 반복 갱신한다.

···
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 = "gradient-descent", 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);

안개 낀 언덕에서 발밑의 기울기를 느끼고 내리막으로 한 걸음씩 내려가는 모습입니다. 각 위치는 모델의 파라미터, 높이는 손실에 대응합니다.

현재 파라미터에서 손실의 기울기를 계산하고 학습률을 곱해 그 반대 방향으로 이동합니다. 전체 데이터로 기울기를 구하면 배치 경사 하강, 일부 표본으로 추정하면 미니배치 방식입니다. 데모는 1차원 볼록 손실의 이동을 단순화합니다.

학습률이 크면 최솟값을 지나치고 작으면 느립니다. 실제 신경망 손실은 여러 차원이고 볼록하지 않을 수 있어 데모처럼 한 번에 최저점이 보이지 않습니다.

언제 쓰나

직접 해를 구하기 어렵거나 큰 모델의 파라미터를 학습할 때. 손실 추이를 봅니다.

페이지로 열기 ↗