python-036Đọc toàn bộ đề miễn phí

Cơ chế Batched Multi-Head Attention Bằng Ký Hiệu Einstein Einsum

ight) V$.

PythonNâng cao35 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ủ đề

numpyeinsumattention mechanismtransformerlinear algebra

Kiến thức tiên quyết: einsum notation, matrix multiplication, attention formula.

Nội dung đề bài

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

  • Sử dụng thành thạo Ký hiệu Tổng Einstein (np.einsum) để biểu diễn các phép nhân ma trận đa chiều phức tạp mà không cần hoán vị chiều (np.transpose).
  • Cài đặt trọn vẹn thuật toán Scaled Dot-Product Attention cốt lõi của mô hình Transformer.
  • Áp dụng mặt nạ Attention Mask với giá trị trừ vô cùng âm (-109) để loại bỏ các token padding.

Mô tả bài toán

Trong kiến trúc Transformer, cơ chế Multi-Head Attention nhận vào ba tensor 4 chiều:

  • Q (Query): có shape (B, H, N, D)
  • K (Key): có shape (B, H, M, D)
  • V (Value): có shape (B, H, M, D_v)

Trong đó: B là batch size, H là số head, N là độ dài chuỗi truy vấn, M là độ dài chuỗi khóa, D là chiều vector head.

Hãy viết hàm: batched_scaled_dot_product_attention(Q: np.ndarray, K: np.ndarray, V: np.ndarray, mask: np.ndarray = None) -> tuple[np.ndarray, np.ndarray]

Quy tắc:

  • Tính ma trận điểm tương đồng (Attention Scores):
  • Sử dụng np.einsum('bhnd,bhmd->bhnm', Q, K).
  • Nhân với hệ số co giãn 1√(D) (với D = Q.shape[-1]).
  • Áp dụng Mask (nếu có):
  • Nếu mask is not None: mask có shape broadcast được với (B, H, N, M) với giá trị 1 (hợp lệ) và 0 (bị che).
  • Tại những vị trí mask == 0, gán giá trị score bằng -1e9.
  • Tính trọng số Attention:
  • Áp dụng Softmax ổn định số học dọc trục cuối cùng (axis=-1).
  • Tính ngữ cảnh đầu ra (Context Output):
  • Nhân trọng số Attention với Value sử dụng np.einsum('bhnm,bhmd->bhnd', attn_weights, V).
  • Kết quả trả về: Tuple (context_output, attn_weights).

Input

  • Các tham số truyền vào hàm/lớp batched_scaled_dot_product_attention hoặc dữ liệu đầu vào theo định dạng mô tả.

Output

  • Kết quả trả về của hàm/lớp batched_scaled_dot_product_attention hoặc dữ liệu in ra màn hình theo đúng đặc tả.

Ràng buộc

  • Thời gian chạy tối đa: 1500ms.
  • Giới hạn bộ nhớ: 512MB.

Ví dụ 1

Input

batched_scaled_dot_product_attention(B=2, H=2, N=3, M=3, D=4, has_mask=False)

Output

True

Giải thích

Hàm được gọi với các tham số mẫu trên và trả về kết quả chính xác theo yêu cầu.

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.