소프트맥스 온도

Softmax temperature

점수를 확률 분포로 바꿀 때 분포가 뾰족할 정도를 조절한다.

···
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 = "softmax-temperature", 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);

여러 후보의 점수를 확률처럼 나눠 줄 때 온도는 선택의 날카로움을 조절하는 손잡이입니다. 낮으면 가장 높은 점수에 몰리고 높으면 후보들이 더 비슷해집니다.

각 logit을 양수 온도 T로 나누고 지수화한 뒤 전체 합으로 나눕니다. 수치 계산에서는 최댓값을 먼저 빼도 확률은 같습니다. 데모는 세 고정 점수에 이 계산을 실제로 적용합니다.

온도 조정만으로 원래 점수가 믿을 만한 확률이 되는 것은 아닙니다. T가 0에 가까우면 거의 최대 점수만 남고, 너무 크면 차이를 구분하기 어렵습니다.

언제 쓰나

분류 출력의 분포나 생성 샘플링의 집중도를 조절할 때. 확률 보정과 구분합니다.

페이지로 열기 ↗