Memahami Konsep Knowledge Distillation: Teori Kompresi Model dan Implementasi dengan PyTorch

Lhuqita Fazry
Deep Learning PyTorch Knowledge Distillation Model Compression
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.

Ilustrasi transfer pengetahuan dari model teacher besar ke model student yang ringan

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.

python
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:

text
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
Machine Learning • Intermediate

Deep Learning Bootcamp

A beginner-friendly, highly interactive bootcamp designed to take you from found...

Daftar

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.

python
# !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:

text
Teacher params: 1,421,066
Student params: 38,282

Rasio 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.

python
# !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:

text
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.

Implementasi training loop knowledge distillation dengan PyTorch

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

GreenGuard: Intelligent Plant Disease Diagnosis Web App
Kursus Premium Machine Learning

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.

Proyek Akhir

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.
7 Weeks Intermediate
Lihat Detail Kursus

Artikel Terkait