- Cho thấy rằng việc thay Dynamic Tanh (DyT) vào vị trí Layer Norm/RMSNorm vốn được dùng như thành phần gần như thiết yếu trong Transformer có thể đạt hiệu năng tương đương hoặc tốt hơn các mô hình chuẩn hóa hiện có
- DyT là một phép toán theo từng phần tử có dạng
DyT(x) = tanh(αx), xuất phát từ quan sát rằng Layer Normalization trong Transformer thường tạo ra ánh xạ đầu vào–đầu ra dạng chữ S tương tự tanh
- Việc triển khai chỉ ở mức vài dòng PyTorch, áp dụng scale và bias lên đầu ra
tanh(alpha * x) bằng các tham số học được alpha, weight, bias
- Đánh giá bao quát toàn bộ các mảng mô hình hóa thị giác, ngôn ngữ, âm thanh và chuỗi DNA, từ ViT, ConvNeXt, MAE, DINO, DiT, LLaMA, wav2vec 2.0, HyenaDNA, Caduceus
- Ngay cả khi không tinh chỉnh hyperparameter riêng, nhiều thiết lập vẫn cho kết quả tương đương hoặc tốt hơn các mô hình đối ứng dựa trên chuẩn hóa, khiến giả định rằng lớp chuẩn hóa là bắt buộc cần được xem xét lại
Điểm mà Dynamic Tanh thay đổi
- DyT là một lớp đơn giản thay thế Layer Norm hoặc RMSNorm trong block Transformer
- Phép toán cốt lõi là
DyT(x) = tanh(αx), được áp dụng theo từng phần tử
- Cho thấy Transformer đã loại bỏ lớp chuẩn hóa vẫn có thể đạt hiệu năng tương đương hoặc cao hơn Transformer chuẩn hóa hiện có
- Ý tưởng bắt nguồn từ quan sát rằng quan hệ đầu vào–đầu ra mà Layer Normalization trong Transformer thường tạo ra giống với hàm scaled tanh
Cách triển khai
- Module DyT có thể được triển khai ngắn gọn trong PyTorch
class DyT(nn.Module):
def __init__(self, num_features, alpha_init_value=0.5):
super().__init__()
self.alpha = nn.Parameter(torch.ones(1) * alpha_init_value)
self.weight = nn.Parameter(torch.ones(num_features))
self.bias = nn.Parameter(torch.zeros(num_features))
def forward(self, x):
x = torch.tanh(self.alpha * x)
return x * self.weight + self.bias
alpha là tham số có thể học được và được đặt giá trị khởi tạo là 0.5
weight và bias cũng là các tham số có thể học được, được áp dụng lên đầu ra tanh(alpha * x)
Quan sát từ Layer Normalization
- Layer Normalization(LN) của Transformer tạo ra ánh xạ đầu vào–đầu ra gần với hàm scaled tanh
- Ở các lớp đầu, ánh xạ này nhìn chung gần tuyến tính
- Càng đi vào các lớp sâu, đường cong chữ S đặc trưng của hàm tanh càng xuất hiện rõ hơn
- Đối tượng quan sát bao gồm Vision Transformer(ViT), mô hình Transformer âm thanh wav2vec 2.0, và các lớp LN được chọn trong Diffusion Transformer(DiT)
Phạm vi đánh giá và kết quả
- DyT được đánh giá trên nhiều kiến trúc và tác vụ
- Thị giác học có giám sát: ViT, ConvNeXt
- Thị giác tự giám sát: MAE, DINO
- Mô hình khuếch tán: DiT
- Mô hình ngôn ngữ lớn: LLaMA
- Âm thanh tự giám sát: wav2vec 2.0
- Mô hình hóa chuỗi DNA: HyenaDNA, Caduceus
- Trong mọi trường hợp, Transformer áp dụng DyT cho thấy hiệu năng tương đương hoặc tốt hơn các mô hình đối ứng dựa trên chuẩn hóa
- Phạm vi đánh giá trải rộng từ nhận thức đến sinh tạo, từ học có giám sát đến tự giám sát, từ thị giác máy tính đến mô hình ngôn ngữ
Tài liệu tham khảo
- Download Paper: bài báo chứa toàn bộ chi tiết của nghiên cứu
- View on GitHub: repository để xem chi tiết triển khai
- View Summary: tóm tắt ngắn gọn kết quả nghiên cứu
Transformers without Normalization được đăng ký là bài báo tại CVPR 2025
1 bình luận
Ý kiến trên Hacker News
Việc điều chỉnh alpha không có mấy tác dụng, nên có thể cần tinh chỉnh siêu tham số đáng kể hoặc khởi tạo tinh vi hơn. Tôi đã thử cả khởi tạo mặc định của PyTorch lẫn khởi tạo trực giao, nhưng không có khác biệt
Hoặc cũng có thể bộ tối ưu hóa scalar tôi dùng không phù hợp. Tôi dùng một bộ tối ưu hóa scalar tùy chỉnh giúp hội tụ nhanh hơn Adam, nhưng với lớp DyT thì nó chỉ có vẻ ngang mức Adam
Cũng có thể đây là kiểu phải sau hàng chục tỷ token mới bắt kịp, nhưng tôi không có ngân sách để thử lâu đến vậy
Nếu có thể thay thế các lớp như vậy, việc này sẽ giúp giảm chi phí tính toán khá đáng kể
tanh chắc chắn cũng sẽ có các ảnh hưởng khác. Vì đôi khi chuẩn hóa đang giải quyết vấn đề điều kiện hóa. Dù vậy, có thêm nhiều lựa chọn thay thế là điều đáng hoan nghênh
Tôi khuyên nên đọc bài báo ResNet gốc của Kaiming He và cộng sự cùng các bài tiếp theo
Với cách tiếp cận hiện đại cho RNN, bài của DeepMind tại https://arxiv.org/abs/2303.06349 đáng để đọc
Điểm cốt lõi là trị riêng lớn nhất, tức bán kính phổ, nên ở gần 1. Điều đó có nghĩa là khi áp dụng lặp lại biến đổi tuyến tính, activation sẽ không tăng lên hoặc nhỏ đi
y = x + f(x)LNinputvớiLNoutputtrong khi sautanh(a*x)vẫn gắn thêm trọng số và biasNếu muốn xem độ tương đồng, chẳng phải nên so với kết quả sau khi bỏ trọng số và bias khỏi đầu ra LayerNorm sao?
Nếu kết quả cuối cùng tốt thì không sao, nhưng nếu tách riêng phần thực sự được thay thế, có lẽ sẽ hiểu rõ hơn chuyện gì đang diễn ra