Memahami Konsep Knowledge Distillation: Teori Kompresi Model dan Implementasi dengan PyTorch
Mengapa Model Besar Sulit Dibawa ke Production
Model deep learning modern berkembang ke arah yang semakin besar. Jumlah parameter pada model vision dan language dengan mudah mencapai ratusan juta. Kapasitas besar tersebut memberikan akurasi tinggi pada tahap riset. Masalah muncul ketika model harus berjalan di production dengan batasan latensi, memori, dan biaya inference.
Kita menghadapi trade-off yang nyata antara akurasi dan efisiensi. Model teacher dengan 50 juta parameter mungkin mencapai akurasi 92% pada CIFAR-10. Model kecil dengan 500 ribu parameter mungkin hanya mencapai 84% bila dilatih secara standar. Selisih 8% tersebut sering membuat tim enggan memakai model kecil.
Knowledge distillation menawarkan jalan tengah yang praktis. Teknik ini melatih model kecil yang disebut student agar meniru perilaku model besar yang disebut teacher. Student tidak hanya belajar dari label keras pada dataset. Student juga belajar dari distribusi probabilitas yang dihasilkan teacher. Informasi tambahan tersebut membantu student mencapai akurasi mendekati teacher dengan biaya inference yang jauh lebih rendah.
Posisi distillation berbeda dari teknik kompresi lain. Pruning memangkas weight yang tidak penting dari arsitektur yang sama. Quantization menurunkan presisi numerik dari float32 ke int8. Distillation melatih arsitektur baru yang memang dirancang kecil sejak awal. Ketiga teknik tersebut dapat dikombinasikan, tetapi distillation memberikan fleksibilitas terbesar dalam memilih arsitektur student.

Gambar: Ilustrasi transfer pengetahuan dari model besar ke model ringan untuk deployment efisien — Sumber: Unsplash
Cara Kerja Soft Target dan Suhu Temperature
Inti knowledge distillation terletak pada konsep soft target. Hard label menyatakan satu kelas bernilai 1 dan kelas lain bernilai 0. Soft target menyatakan probabilitas untuk setiap kelas, misalnya anjing 0.75, kucing 0.20, dan kelinci 0.05. Distribusi tersebut mengandung pengetahuan tentang kemiripan antar kelas yang dipelajari teacher.
Suhu temperature dengan simbol T mengontrol kelembutan distribusi tersebut. Softmax standar menggunakan T=1 dan menghasilkan distribusi yang tajam. Nilai T yang lebih besar seperti 4 atau 10 membuat distribusi lebih datar. Kelas yang semula hampir nol menjadi terlihat. Informasi tentang kelas kedua dan ketiga yang menurut teacher cukup mirip ikut tersampaikan ke student.
Fungsi loss pada distillation menggabungkan dua komponen. Komponen pertama adalah KLDivLoss antara output student dan soft target teacher pada suhu T. Komponen kedua adalah CrossEntropyLoss antara output student dan hard label pada suhu 1. Parameter alpha mengatur bobot keduanya. Nilai alpha=0.7 berarti 70% bobot pada pengetahuan teacher dan 30% pada label asli.
import torch
import torch.nn.functional as F
# Logits dummy dari teacher untuk 3 kelas
teacher_logits = torch.tensor([[3.0, 1.0, 0.2]])
for T in [1, 4, 10]:
soft = F.softmax(teacher_logits / T, dim=1)
print(f"T={T}: {soft.numpy().round(3)}")Penjelasan kode di atas berfokus pada alur logika yang ingin kita tunjukkan. Kita membuat satu vektor logits dari teacher. Kita membaginya dengan suhu berbeda sebelum melewatkannya ke fungsi softmax. Hasil yang diharapkan menunjukkan distribusi semakin datar ketika T naik.
Output:
T=1: [[0.836 0.113 0.051]]
T=4: [[0.475 0.288 0.236]]
T=10: [[0.388 0.318 0.294]]Perbedaan tersebut menjelaskan mengapa temperature penting. Pada T=1, student hampir hanya melihat kelas pertama. Pada T=4, student melihat struktur kemiripan antar kelas. Pada T=10, distribusi menjadi terlalu datar dan sinyal diskriminatif melemah. Praktik umum memakai rentang 2 sampai 10.
Deep Learning Bootcamp
A beginner-friendly, highly interactive bootcamp designed to take you from found...
Membangun Teacher dan Student Model dengan PyTorch
Tahap berikutnya adalah menyiapkan dua arsitektur dengan kapasitas berbeda. Teacher memakai CNN yang lebih dalam dengan jumlah channel besar. Student memakai CNN ringan dengan channel kecil dan layer lebih sedikit. Perbedaan jumlah parameter bisa mencapai 10 sampai 20 kali lipat.
Alur training mengikuti pola yang konsisten. Teacher dilatih terlebih dahulu sampai konvergen, lalu dibekukan dalam mode evaluasi. Student dilatih dari awal dengan distillation loss. Selama training student, gradient hanya mengalir ke parameter student. Parameter teacher tidak diperbarui.
Kita memakai dataset Fashion-MNIST pada contoh ini. Dataset tersebut berukuran kecil sehingga runnable di sandbox CPU. Prinsip yang sama berlaku untuk CIFAR-10 atau dataset custom dengan mengganti data loader saja.
# !pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu
import torch
import torch.nn as nn
import torch.nn.functional as F
class TeacherCNN(nn.Module):
def __init__(self):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(1, 64, 3, padding=1),
nn.ReLU(),
nn.Conv2d(64, 128, 3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(128, 256, 3, padding=1),
nn.ReLU(),
nn.AdaptiveAvgPool2d((4, 4))
)
self.classifier = nn.Sequential(
nn.Flatten(),
nn.Linear(256 * 4 * 4, 256),
nn.ReLU(),
nn.Linear(256, 10)
)
def forward(self, x):
return self.classifier(self.features(x))
class StudentCNN(nn.Module):
def __init__(self):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(1, 16, 3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(16, 32, 3, padding=1),
nn.ReLU(),
nn.AdaptiveAvgPool2d((4, 4))
)
self.classifier = nn.Sequential(
nn.Flatten(),
nn.Linear(32 * 4 * 4, 64),
nn.ReLU(),
nn.Linear(64, 10)
)
def forward(self, x):
return self.classifier(self.features(x))
def distillation_loss(student_logits, teacher_logits, targets, T=4.0, alpha=0.7):
soft_teacher = F.softmax(teacher_logits / T, dim=1)
soft_student = F.log_softmax(student_logits / T, dim=1)
distill = F.kl_div(soft_student, soft_teacher, reduction="batchmean") * (T * T)
ce = F.cross_entropy(student_logits, targets)
return alpha * distill + (1 - alpha) * ce
teacher = TeacherCNN()
student = StudentCNN()
n_teacher = sum(p.numel() for p in teacher.parameters())
n_student = sum(p.numel() for p in student.parameters())
print(f"Teacher params: {n_teacher:,}")
print(f"Student params: {n_student:,}")Penjelasan kode di atas menekankan workflow, bukan detail baris per baris. Kita mendefinisikan dua arsitektur dengan pola yang sama tetapi kapasitas berbeda. Kita menulis fungsi loss yang menggabungkan pengetahuan teacher dan label asli. Output yang diharapkan menampilkan perbandingan jumlah parameter sebagai bukti kompresi.
Output:
Teacher params: 1,421,066
Student params: 38,282Rasio sekitar 37 banding 1 tersebut menunjukkan potensi penghematan memori dan percepatan inference. Student hanya membawa sekitar 2.7% parameter teacher.
Melatih Student dengan Distillation Loss sampai Evaluasi Akurasi
Training student memakai loop standar PyTorch dengan satu perbedaan penting. Setiap batch melewati teacher tanpa gradient untuk menghasilkan soft target. Logits teacher dan student kemudian masuk ke fungsi distillation loss. Optimizer Adam memperbarui parameter student berdasarkan loss gabungan tersebut.
Kita membandingkan tiga skenario untuk mengukur manfaat distillation. Skenario pertama adalah teacher sebagai batas atas akurasi. Skenario kedua adalah student yang dilatih hanya dengan hard label sebagai baseline. Skenario ketiga adalah student yang dilatih dengan distillation. Perbandingan yang adil memakai arsitektur, optimizer, dan jumlah epoch yang sama untuk kedua student.
# !pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu
import torch
import torch.nn.functional as F
from torchvision import datasets, transforms
from torch.utils.data import DataLoader, Subset
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
full_train = datasets.FashionMNIST(root="./data", train=True, download=True, transform=transform)
full_test = datasets.FashionMNIST(root="./data", train=False, download=True, transform=transform)
# Subset kecil agar cepat di CPU
train_loader = DataLoader(Subset(full_train, range(5000)), batch_size=128, shuffle=True)
test_loader = DataLoader(Subset(full_test, range(1000)), batch_size=256)
device = torch.device("cpu")
teacher = TeacherCNN().to(device)
student = StudentCNN().to(device)
optimizer = torch.optim.Adam(student.parameters(), lr=1e-3)
teacher.eval()
for p in teacher.parameters():
p.requires_grad = False
T, alpha = 4.0, 0.7
for epoch in range(3):
student.train()
total = 0.0
for x, y in train_loader:
x, y = x.to(device), y.to(device)
with torch.no_grad():
t_logits = teacher(x)
s_logits = student(x)
loss = distillation_loss(s_logits, t_logits, y, T=T, alpha=alpha)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total += loss.item()
print(f"Epoch {epoch+1}: loss={total/len(train_loader):.4f}")
student.eval()
correct = 0
total_n = 0
with torch.no_grad():
for x, y in test_loader:
pred = student(x.to(device)).argmax(dim=1)
correct += (pred == y).sum().item()
total_n += y.size(0)
print(f"Student distilled accuracy: {100*correct/total_n:.2f}%")Penjelasan kode di atas berfokus pada alasan setiap langkah. Kita membekukan teacher agar berfungsi sebagai pemberi referensi yang stabil. Kita memakai subset data agar eksperimen selesai dalam hitungan menit di CPU. Output yang diharapkan berupa penurunan loss per epoch dan skor akurasi akhir.
Output:
Epoch 1: loss=0.6246
Epoch 2: loss=0.5560
Epoch 3: loss=0.5249
Student distilled accuracy: 72.50%Angka di atas berasal dari run singkat tiga epoch. Pola umum pada eksperimen penuh menunjukkan student distilled unggul 3 sampai 6 poin dibanding student baseline dengan arsitektur identik. Teacher tetap sedikit lebih tinggi, tetapi gap menyempit signifikan dibanding training standar.

Gambar: Implementasi training loop distillation dengan PyTorch pada lingkungan pengembangan — Sumber: Unsplash
Memilih Temperature dan Alpha yang Tepat untuk Kasus Nyata
Pemilihan T dan alpha menentukan keberhasilan distillation. Nilai T antara 2 dan 5 cocok untuk dataset dengan kelas yang cukup berbeda seperti Fashion-MNIST. Nilai T antara 5 dan 10 membantu pada dataset dengan banyak kelas mirip seperti CIFAR-100. Nilai alpha antara 0.5 dan 0.9 memberi porsi besar pada pengetahuan teacher ketika teacher memang akurat.
Distillation dapat gagal pada beberapa kondisi. Gap kapasitas yang terlalu besar membuat student tidak mampu meniru teacher. Teacher yang overfit memberikan soft target yang terlalu tajam sehingga mirip hard label. Dataset yang terlalu kecil membuat student menghafal soft target tanpa generalisasi. Mitigasi praktis meliputi penggunaan teacher intermediate, augmentasi data, dan early stopping berbasis akurasi validasi.
Checklist sebelum deployment membantu memastikan manfaat nyata. Kita mengukur akurasi test, jumlah parameter, ukuran file model, dan latensi rata-rata per batch. Kita membandingkan student distilled dengan baseline kuantisasi. Kita memvalidasi kalibrasi probabilitas bila model dipakai untuk keputusan berisiko. Kita juga mencatat konfigurasi T dan alpha terbaik agar eksperimen mudah direproduksi oleh tim. Langkah tersebut memastikan kompresi tidak mengorbankan keandalan.
Ingin mendalami kompresi model dan deployment deep learning sampai production? Pelajari kurikulum Deep Learning di Rumah Coding dan bangun portofolio model efisien yang siap production.
Kursus Terkait
Deep Learning Bootcamp
A beginner-friendly, highly interactive bootcamp designed to take you from foundational concepts to deploying real-world Artificial Intelligence applications. Through a completely project-based approach, you will master the core of Deep Learning, Artificial Neural Networks, and Computer Vision using Python and TensorFlow, ultimately building a professional-grade AI web application for your portfolio.
GreenGuard: Intelligent Plant Disease Diagnosis Web App
- Interactive Image Upload UI: A clean, user-friendly interface built with Streamlit that supports drag-and-drop image uploads directly from a computer or mobile phone.
- Real-Time AI Inference: Utilizes a lightweight, optimized CNN model (like MobileNetV2) to process the image and return a diagnosis in seconds without heavy server load.
- Confidence Scoring Dashboard: Visually displays the model's prediction probability (e.g., "95% confident this is Tomato Late Blight") using interactive progress bars or charts.