ai-030Đọc toàn bộ đề miễn phí

Triển khai Cơ chế Self-Attention (Scaled Dot-Product Attention) từ Gốc

Attention(Q, K, V) = softmax(Q KT√(dk) + M) V

AINâng cao45 phút

Tiến độ của tôi ở bài này

Điểm được lưu vào tài khoản sau khi chấm bài.

Đang tải điểm của bạn…

Kiến thức và chủ đề

deep-learningtransformerattentionscaled-dot-productcausal-masksoftmaxpytorch

Kiến thức tiên quyết: stable-softmax-logits, pytorch-tensor-basics-cpu.

Nội dung đề bài

Mục tiêu kiến thức

  • Công thức nền tảng trong bài báo kinh điển *"Attention Is All You Need"* (Vaswani et al., 2017):

Attention(Q, K, V) = softmax(Q KT√(dk) + M) V Trong đó:

  • Q (Query), K (Key), V (Value) có số chiều chiều cuối dk.
  • Tỷ lệ co 1√(dk) giúp tránh việc tích vô hướng quá lớn làm gradient của Softmax bị triệt tiêu (vanishing gradient).
  • M là ma trận mặt nạ chú ý (Attention Mask / Causal Mask): Tại các vị trí bị che (ví dụ token tương lai trong mô hình sinh GPT), giá trị được gán bằng hằng số rất âm (ví dụ -109) để sau khi qua Softmax, xác suất chú ý bằng đúng 0.

Yêu cầu

Viết hàm scaled_dot_product_attention(Q: torch.Tensor, K: torch.Tensor, V: torch.Tensor, mask: torch.Tensor | None = None) -> tuple[torch.Tensor, torch.Tensor]:

  • Đầu vào: Các tensor có kích thước (B, H, Sq, D) hoặc (B, Sq, D).
  • dk = Q.size(-1).
  • Tính ma trận điểm tương đồng (Attention Scores):

scores = Q · KT√(dk) (Lưu ý chuyển vị 2 chiều cuối cùng của K).

  • Nếu có mask: Áp dụng mặt nạ:
  • Giả định mask là tensor nhị phân hoặc boolean trong đó giá trị 0 (hoặc False) biểu thị vị trí cần che. Thay thế các vị trí đó trong scores bằng -109.
  • Tính trọng số chú ý:

attn_weights = torch.softmax(scores, dim=-1)

  • Nhân với Value:

output = attn_weights · V

  • Trả về tuple (output, attn_weights).

Input

  • Hàm scaled_dot_product_attention(Q, K, V, mask): Các tham số đầu vào chứa dữ liệu Tensor/mảng NumPy hoặc giá trị siêu tham số tương ứng.

Output

  • Hàm scaled_dot_product_attention: Trả về kết quả kiểu tuple[torch.Tensor, torch.Tensor] theo đúng đặc tả kỹ thuật và kích thước quy định.

Ràng buộc

  • Thời gian chạy tối đa: 3000ms.
  • Giới hạn bộ nhớ: 512MB.
  • Dữ liệu đầu vào hợp lệ theo đúng kiểu dữ liệu và miền giá trị được mô tả.

Ví dụ 1

Input

B, H, S, D = (2, 4, 8, 16)
Q = torch.randn(B, H, S, D)
K = torch.randn(B, H, S, D)
V = torch.randn(B, H, S, D)
out, weights = scaled_dot_product_attention(Q, K, V)
sum_weights = torch.sum(weights, dim=-1)

Output

Tensor shape: (2, 4, 8), dtype=torch.float32

Giải thích

Hàm/lớp được gọi với các tham số mẫu trên và trả về kết quả số học / kích thước tensor tương ứng theo đúng thiết kế.

Ví dụ 2

Input

S, D = (3, 4)
Q = torch.randn(1, S, D)
K = torch.randn(1, S, D)
V = torch.randn(1, S, D)
causal_mask = torch.tril(torch.ones(S, S))
out, weights = scaled_dot_product_attention(Q, K, V, mask=causal_mask)

Output

(Tensor shape: (1, 3, 4), dtype=torch.float32, Tensor shape: (1, 3, 3), dtype=torch.float32)

Giải thích

Hàm/lớp được gọi với các tham số mẫu trên và trả về kết quả số học / kích thước tensor tương ứng theo đúng thiết kế.

3 cấp độ gợi ýMở dần khi bạn thật sự cần hỗ trợ.
Phân tích lời giảiGiải thích hướng tư duy và thuật toán.
Code tham khảoDùng để đối chiếu sau khi tự làm.

Gợi ý và lời giải chỉ mở sau khi bạn bấm Nộp bài. Giáo viên và quản trị viên mở được ngay.