Phân tích thuật toán đánh giá thẩm mỹ dựa trên mạng MLP

Giới thiệu về mạng MLP và ứng dụng trong dự đoán điểm thẩm mỹ ảnh

Mạng improved-aesthetic-predictor là một công cụ đánh giá chất lượng thẩm mỹ của ảnh dựa trên kiến trúc CLIP và MLP. Công cụ này có khả năng tự động đánh giá mức độ thẩm mỹ của hình ảnh, kết hợp giữa mô hình CLIP trích xuất đặc trưng và mạng MLP thực hiện ánh xạ đặc trưng thành điểm số.

Mạng MLP là gì?

Mạng Perceptron đa lớp (MLP) là một loại mạng nơ-ron cơ bản bao gồm ba lớp chính: lớp đầu vào, lớp ẩn và lớp đầu ra. Trong dự án improved-aesthetic-predictor, MLP được sử dụng để chuyển đổi các đặc trưng hình ảnh được trích xuất từ mô hình CLIP thành điểm thẩm mỹ cụ thể.

Cấu trúc mạng MLP trong dự án

Lớp MLP được định nghĩa trong file train_predictor.py với cấu trúc như sau:

class MLP(pl.LightningModule):
    def __init__(self, input_dim, input_col='emb', target_col='avg_rating'):
        super().__init__()
        self.input_dim = input_dim
        self.net = nn.Sequential(
            nn.Linear(self.input_dim, 1024),
            nn.ReLU(),
            nn.Dropout(0.2),
            nn.Linear(1024, 128),
            nn.ReLU(),
            nn.Dropout(0.2),
            nn.Linear(128, 64),
            nn.ReLU(),
            nn.Dropout(0.1),
            nn.Linear(64, 1)
        )

Cấu trúc trên bao gồm 3 lớp ẩn với hàm kích hoạt ReLU, các lớp Dropout giúp giảm hiện tượng overfitting. Đầu vào là vector đặc trưng 768 chiều từ CLIP, đầu ra là giá trị điểm thẩm mỹ đơn.

Cơ chế hoạt động của CLIP + MLP

  1. Trích xuất đặc trưng: Dùng mô hình CLIP (ví dụ: ViT-L/14) để trích xuất vector 768 chiều từ hình ảnh đầu vào.
  2. Chuyển đổi đặc trưng: Mạng MLP ánh xạ vector đặc trưng thành điểm số thẩm mỹ.
  3. Tối ưu mô hình: Sử dụng hàm mất mát MSE (Mean Squared Error) để huấn luyện mô hình.

Quy trình huấn luyện

Hàm huấn luyện được định nghĩa như sau:

def training_step(self, batch, idx):
    features = batch[self.input_col]
    labels = batch[self.target_col].view(-1, 1)
    predictions = self.net(features)
    loss = F.mse_loss(predictions, labels)
    self.log("train_loss", loss)
    return loss

Mô hình được huấn luyện bằng thuật toán tối ưu Adam để giảm thiểu sai số giữa điểm dự đoán và điểm đánh giá thực tế từ dữ liệu huấn luyện.

Các mô hình đã huấn luyện sẵn

Dự án cung cấp một số file mô hình đã huấn luyện sẵn:

  • ava+logos-l14-linearMSE.pth: Mô hình tuyến tính huấn luyện trên tập AVA và logos.
  • ava+logos-l14-reluMSE.pth: Mô hình có lớp kích hoạt ReLU.
  • sac+logos+ava1-l14-linearMSE.pth: Mô hình tổng hợp từ nhiều tập dữ liệu khác nhau.

Hướng dẫn sử dụng nhanh

Cài đặt môi trường

git clone https://gitcode.com/gh_mirrors/im/improved-aesthetic-predictor
cd improved-aesthetic-predictor

Đánh giá điểm thẩm mỹ cho ảnh

Bạn có thể sử dụng script simple_inference.py để thực hiện đánh giá:

# Ví dụ mã nguồn đơn giản
model = MLP(768)
model.load_state_dict(torch.load("sac+logos+ava1-l14-linearMSE.pth"))
image = preprocess(Image.open("test.jpg")).unsqueeze(0)
with torch.no_grad():
    features = clip_model.encode_image(image)
    score = model(features)
    print(f"Điểm thẩm mỹ: {score.item()}")

Đoạn mã trên cho thấy cách tải mô hình, trích xuất đặc trưng từ ảnh và dự đoán điểm thẩm mỹ.

Thẻ: CLIP mlp PyTorch Aesthetic Score Image Evaluation

Đăng vào ngày 28 tháng 9 lúc 09:48