Chuyển đến nội dung chính

Bài 10: Vision Transformer (ViT) & CLIP

ViT: Transformer cho ảnh — patch embedding, position encoding, self-attention trên ảnh. CLIP: kết nối text và image trong cùng embedding space. Zero-shot image classification. Ứng dụng CLIP trong search, recommendation.

🧠 AI & ML — Bài 9 Bài 10: Vision Transformer (ViT) & CLIP

Computer Vision với Deep Learning: Từ CNN đến Vision Transformer

Phần 3: Segmentation & Modern CV

xdev.asia

Giới thiệu

Transformer đã thay đổi NLP (GPT, BERT). Giờ nó thay đổi cả Computer Vision. ViT (Vision Transformer) chứng minh: Transformer có thể vượt CNN trên image classification. CLIP (OpenAI) kết nối text và image trong cùng không gian → zero-shot classification, image search, multimodal AI.

🎯 ViT + CLIP là nền tảng cho multimodal AI hiện đại: GPT-4o Vision, Gemini, Claude Vision.


1. Vision Transformer (ViT)

1.1 Ý tưởng: Ảnh → Sequence of Patches

Input Image (224×224)
    ↓
Chia thành patches (16×16 pixels mỗi patch)
→ 14 × 14 = 196 patches
    ↓
Mỗi patch → flatten → linear projection → patch embedding
→ 196 vectors, mỗi vector dimension D (ví dụ 768)
    ↓
Thêm [CLS] token + Position Embeddings
→ 197 tokens
    ↓
Transformer Encoder (12 layers)
→ Self-attention trên tất cả patches
    ↓
[CLS] token output → Classification Head
→ Prediction

1.2 Patch Embedding

"""ViT: biến ảnh thành sequence of patches"""
import torch
import torch.nn as nn

class PatchEmbedding(nn.Module):
    def __init__(self, img_size=224, patch_size=16, in_channels=3, embed_dim=768):
        super().__init__()
        self.num_patches = (img_size // patch_size) ** 2  # 196
        self.patch_size = patch_size

        # Linear projection of flattened patches
        # Dùng Conv2d stride=patch_size → tương đương cắt patches + linear
        self.projection = nn.Conv2d(
            in_channels, embed_dim,
            kernel_size=patch_size,
            stride=patch_size
        )

    def forward(self, x):
        # x: (B, 3, 224, 224)
        x = self.projection(x)           # (B, 768, 14, 14)
        x = x.flatten(2)                 # (B, 768, 196)
        x = x.transpose(1, 2)           # (B, 196, 768) = sequence of patches
        return x

# Test
patch_embed = PatchEmbedding()
img = torch.randn(1, 3, 224, 224)
patches = patch_embed(img)
print(f"Patches: {patches.shape}")  # (1, 196, 768)

1.3 ViT Architecture

"""Simplified ViT implementation"""
class VisionTransformer(nn.Module):
    def __init__(self, img_size=224, patch_size=16, in_channels=3,
                 embed_dim=768, num_heads=12, num_layers=12, num_classes=1000):
        super().__init__()

        # Patch Embedding
        self.patch_embed = PatchEmbedding(img_size, patch_size, in_channels, embed_dim)
        num_patches = self.patch_embed.num_patches

        # [CLS] token — learnable
        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))

        # Position Embeddings — learnable
        self.pos_embed = nn.Parameter(torch.randn(1, num_patches + 1, embed_dim))

        # Transformer Encoder
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embed_dim,
            nhead=num_heads,
            dim_feedforward=embed_dim * 4,
            activation="gelu",
            batch_first=True,
        )
        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)

        # Classification Head
        self.norm = nn.LayerNorm(embed_dim)
        self.head = nn.Linear(embed_dim, num_classes)

    def forward(self, x):
        B = x.shape[0]

        # Patch embed
        x = self.patch_embed(x)           # (B, 196, 768)

        # Prepend [CLS] token
        cls_tokens = self.cls_token.expand(B, -1, -1)  # (B, 1, 768)
        x = torch.cat([cls_tokens, x], dim=1)          # (B, 197, 768)

        # Add position embeddings
        x = x + self.pos_embed                          # (B, 197, 768)

        # Transformer
        x = self.transformer(x)                         # (B, 197, 768)

        # Classify from [CLS] token
        cls_output = self.norm(x[:, 0])                 # (B, 768)
        logits = self.head(cls_output)                  # (B, num_classes)

        return logits

model = VisionTransformer(num_classes=10)
img = torch.randn(2, 3, 224, 224)
out = model(img)
print(f"Output: {out.shape}")  # (2, 10)

1.4 ViT vs CNN

CNN (ResNet)Transformer (ViT)
Inductive biasTranslation invariance, localityÍt bias → cần nhiều data hơn
Data cầnÍt (<1M ảnh)Nhiều (>10M ảnh)
ScalabilityĐỉnh ở ~150 layersScale tốt: ViT-G (2B params)
Global contextCần nhiều layersSelf-attention ngay (mọi patch nói chuyện)
PretrainedImageNet (1.3M)JFT-300M, LAION-5B

2. Sử dụng ViT Pretrained

"""ViT pretrained — image classification"""
import torch
from transformers import ViTForImageClassification, ViTImageProcessor
from PIL import Image

# Load model + processor
processor = ViTImageProcessor.from_pretrained("google/vit-base-patch16-224")
model = ViTForImageClassification.from_pretrained("google/vit-base-patch16-224")

# Inference
image = Image.open("cat.jpg")
inputs = processor(images=image, return_tensors="pt")

with torch.no_grad():
    outputs = model(**inputs)

# Top-5 predictions
logits = outputs.logits
probs = torch.softmax(logits, dim=-1)
top5_prob, top5_idx = probs.topk(5)

for prob, idx in zip(top5_prob[0], top5_idx[0]):
    label = model.config.id2label[idx.item()]
    print(f"  {label:30s} {prob.item():.4f}")

2.1 Fine-tune ViT cho Custom Dataset

"""Fine-tune ViT trên dataset riêng"""
from transformers import ViTForImageClassification, TrainingArguments, Trainer
from datasets import load_dataset
import torch

# Load dataset
dataset = load_dataset("food101", split={"train": "train[:5000]", "test": "validation[:1000]"})

# Processor
processor = ViTImageProcessor.from_pretrained("google/vit-base-patch16-224")

def transform(examples):
    examples["pixel_values"] = [
        processor(images=img, return_tensors="pt")["pixel_values"][0]
        for img in examples["image"]
    ]
    return examples

dataset = dataset.with_transform(transform)

# Model
model = ViTForImageClassification.from_pretrained(
    "google/vit-base-patch16-224",
    num_labels=101,
    ignore_mismatched_sizes=True,
)

# Training
training_args = TrainingArguments(
    output_dir="./vit-food101",
    num_train_epochs=5,
    per_device_train_batch_size=16,
    learning_rate=2e-5,
    evaluation_strategy="epoch",
    save_strategy="epoch",
    load_best_model_at_end=True,
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=dataset["train"],
    eval_dataset=dataset["test"],
)

trainer.train()

3. CLIP — Kết nối Text và Image

3.1 CLIP hoạt động thế nào?

Training (400M image-text pairs từ internet):

Image Encoder (ViT)     Text Encoder (Transformer)
"cat.jpg"   → [0.8, -0.2, ...]    "a photo of a cat" → [0.7, -0.1, ...]
                    ↕ Cosine Similarity = HIGH ✅

"cat.jpg"   → [0.8, -0.2, ...]    "a photo of a car" → [-0.3, 0.6, ...]
                    ↕ Cosine Similarity = LOW ❌

→ Learn: image và text cùng ý nghĩa = vectors gần nhau

3.2 Zero-shot Image Classification

"""CLIP: phân loại ảnh MÀ KHÔNG CẦN TRAIN!"""
import torch
from PIL import Image
from transformers import CLIPProcessor, CLIPModel

# Load CLIP
model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32")
processor = CLIPProcessor.from_pretrained("openai/clip-vit-base-patch32")

# Image + text labels (bạn tự define!)
image = Image.open("animal.jpg")

# Candidate labels — BẤT KỲ text nào!
labels = [
    "a photo of a cat",
    "a photo of a dog",
    "a photo of a bird",
    "a photo of a fish",
    "a photo of a horse",
]

# Encode
inputs = processor(
    text=labels,
    images=image,
    return_tensors="pt",
    padding=True,
)

with torch.no_grad():
    outputs = model(**inputs)

# Similarities
logits_per_image = outputs.logits_per_image  # (1, 5)
probs = logits_per_image.softmax(dim=1)

for label, prob in zip(labels, probs[0]):
    print(f"  {label:30s} {prob.item():.4f}")

3.3 Image Search với CLIP

"""CLIP: tìm ảnh bằng text query"""
import os
import torch
import numpy as np
from PIL import Image
from transformers import CLIPProcessor, CLIPModel

model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32")
processor = CLIPProcessor.from_pretrained("openai/clip-vit-base-patch32")

# 1. Encode tất cả ảnh trong thư mục (1 lần)
image_dir = "my_photos/"
image_paths = [os.path.join(image_dir, f) for f in os.listdir(image_dir)
               if f.endswith(('.jpg', '.png'))]

image_embeddings = []
for path in image_paths:
    img = Image.open(path)
    inputs = processor(images=img, return_tensors="pt")
    with torch.no_grad():
        emb = model.get_image_features(**inputs)
    emb = emb / emb.norm(dim=-1, keepdim=True)  # Normalize
    image_embeddings.append(emb)

image_embeddings = torch.cat(image_embeddings)  # (N, 512)

# 2. Search bằng text query
def search_images(query, top_k=5):
    inputs = processor(text=[query], return_tensors="pt")
    with torch.no_grad():
        text_emb = model.get_text_features(**inputs)
    text_emb = text_emb / text_emb.norm(dim=-1, keepdim=True)

    # Cosine similarity
    similarities = (text_emb @ image_embeddings.T).squeeze()
    top_indices = similarities.argsort(descending=True)[:top_k]

    print(f"🔍 Query: '{query}'")
    for i, idx in enumerate(top_indices):
        print(f"  {i+1}. {image_paths[idx]} (score: {similarities[idx]:.3f})")

# Test
search_images("sunset at the beach")
search_images("a person cooking food")
search_images("mountains with snow")

3.4 CLIP Applications

🔍 Image Search:        Text query → tìm ảnh tương tự
🏷️ Auto Tagging:        Tự gán tags cho ảnh (zero-shot)
🛡️ Content Moderation:  "NSFW content" → filter
📊 Image Clustering:    Nhóm ảnh theo semantic meaning
🎯 Recommendation:      "Show me photos similar to this style"
📝 Image Captioning:    CLIP + GPT → describe ảnh

4. DINOv2 — Self-supervised ViT

"""DINOv2: ViT pretrained không cần labels — Meta AI"""
import torch
from transformers import AutoModel, AutoImageProcessor
from PIL import Image

processor = AutoImageProcessor.from_pretrained("facebook/dinov2-base")
model = AutoModel.from_pretrained("facebook/dinov2-base")

# Extract features
image = Image.open("product.jpg")
inputs = processor(images=image, return_tensors="pt")

with torch.no_grad():
    outputs = model(**inputs)

# CLS token embedding — dùng cho classification, retrieval
cls_embedding = outputs.last_hidden_state[:, 0]  # (1, 768)

# Patch tokens — dùng cho segmentation, dense prediction
patch_embeddings = outputs.last_hidden_state[:, 1:]  # (1, 196, 768)

print(f"CLS embedding: {cls_embedding.shape}")
print(f"Patch embeddings: {patch_embeddings.shape}")

Tóm tắt

ConceptGhi nhớ
ViTTransformer cho ảnh: chia patches → self-attention
Patch EmbeddingẢnh 224×224 → 196 patches (16×16 each) → embed
CLIPKết nối text + image trong cùng embedding space
Zero-shotClassify ảnh KHÔNG cần train — chỉ cần text labels
Image SearchEncode ảnh + text → cosine similarity → tìm kiếm
DINOv2Self-supervised ViT, features universal

Bài tập tổng hợp

  1. ViT Classification: Dùng ViT pretrained classify 20 ảnh. So sánh accuracy với ResNet-50.
  2. CLIP Zero-shot: Tạo 10 custom labels tiếng Việt. CLIP classify ảnh đúng không?
  3. Image Search: Xây image search engine cho 100 ảnh cá nhân. Search bằng tiếng Việt.
  4. CLIP Similarity: 2 ảnh bất kỳ → CLIP similarity score. Nào giống nhau nhất?

Bài tiếp theo: OCR & Document Understanding — nhận dạng chữ và hiểu tài liệu.