Thủ thuật
Tối ưu hóa GPU Kernel với TileLang: Từ GEMM đến FlashAttention
(giờ Việt Nam)
Tóm tắt AI
TileLang là ngôn ngữ Python giúp đơn giản hóa việc thiết kế GPU kernel hiệu năng cao, cho phép tự động hóa các tác vụ phức tạp như Tensor-Core GEMM và FlashAttention mà không cần can thiệp sâu vào CUDA.
Bản dịch AI

Trong hướng dẫn này, chúng ta sẽ khám phá TileLang như một ngôn ngữ đặc thù miền (domain-specific language) dựa trên Python cấp cao, dùng để thiết kế và biên dịch các GPU kernel tối ưu hiệu năng thông qua TVM. Chúng ta bắt đầu bằng việc xác thực môi trường CUDA và thiết lập các tiện ích đo kiểm (benchmarking) và kiểm chứng số học có thể tái sử dụng, sau đó triển khai dần dần các tác vụ: cộng vector, nhân ma trận tensor-core dạng tile, khám phá lịch trình (schedule exploration), các epilogue GEMM được hợp nhất (fused), softmax theo hàng và FlashAttention. Trong suốt hướng dẫn, chúng ta làm việc trực tiếp với các tile bộ nhớ chia sẻ (shared-memory tiles), các mảnh thanh ghi (register fragments), vòng lặp đường ống (pipelined loops), các nguyên hàm lặp song song, các phép khử (reductions) và toán tử GEMM tensor-core, đồng thời để trình biên dịch quản lý việc ánh xạ luồng, bố cục bộ nhớ, đồng bộ hóa, vector hóa và tạo mã lệnh CUDA cấp thấp. Chúng ta cũng so sánh các kernel của mình với các baseline từ PyTorch và cuBLAS, kiểm tra mã nguồn CUDA được tạo ra, đánh giá thông lượng bộ nhớ và tính toán, đồng thời sử dụng tính năng tự động tinh chỉnh (autotuning) để xác định các cấu hình kernel phụ thuộc vào kiến trúc.
Chúng ta cấu hình môi trường Google Colab CUDA, cài đặt TileLang với tùy chọn dự phòng nightly, và import các module PyTorch và TileLang cần thiết. Chúng ta định nghĩa các tiện ích đo kiểm, xác thực và báo cáo có thể tái sử dụng để đo độ trễ của kernel và so sánh kết quả đầu ra bằng sai số tương đối. Sau đó, chúng ta triển khai một kernel cộng vector bằng TileLang, thực thi trên GPU, so sánh băng thông của nó với PyTorch và kiểm tra mã nguồn CUDA do trình biên dịch tạo ra.
Chúng ta triển khai một kernel nhân ma trận tensor-core dạng tile, di chuyển các tile đầu vào qua bộ nhớ toàn cục (global memory), bộ nhớ chia sẻ và các mảnh thanh ghi. Chúng ta kiểm soát kích thước tile, các giai đoạn đường ống, số lượng luồng và L2 swizzling trong khi cho phép TileLang tạo ra các lệnh tensor-core, đồng bộ hóa và logic truyền dữ liệu bộ nhớ. Sau đó, chúng ta đo kiểm một số cấu hình lịch trình, xác minh độ chính xác số học của chúng và xác định cấu hình kernel có hiệu năng cao nhất phụ thuộc vào kiến trúc.
Chúng ta mở rộng kernel nhân ma trận bằng cách hợp nhất phép cộng bias và hàm kích hoạt GELU trực tiếp vào bộ tích lũy nằm trong thanh ghi (register-resident accumulator). Chúng ta giảm lưu lượng truy cập bộ nhớ toàn cục trung gian bằng cách hoàn thành phần epilogue trước khi ghi tensor đầu ra cuối cùng và so sánh bản triển khai hợp nhất này với việc thực thi eager của PyTorch. Chúng ta cũng triển khai một kernel softmax theo hàng bằng cách sử dụng các phép khử cực đại và tổng ở cấp độ mảnh (fragment-level), đồng thời giữ cho quá trình chuẩn hóa chủ yếu nằm trong các thanh ghi.
Chúng ta triển khai một kernel FlashAttention forward hợp nhất, xử lý các tile query, key và value mà không cần hiện thực hóa toàn bộ ma trận điểm chú ý (attention-score matrix) trong bộ nhớ toàn cục. Chúng ta áp dụng các bản cập nhật softmax trực tuyến (online softmax) bằng cách sử dụng các giá trị cực đại chạy, tổng chuẩn hóa, hệ số tái tỷ lệ và các phép nhân ma trận tensor-core dạng tile. Chúng ta xác thực cả cơ chế chú ý nhân quả (causal) và phi nhân quả (non-causal) so với cơ chế scaled dot-product attention của PyTorch và so sánh độ trễ cũng như thông lượng tính toán của chúng.
Chúng ta định nghĩa một không gian tìm kiếm tự động tinh chỉnh (autotuning search space) bao gồm kích thước tile ma trận, kích thước khối K, độ sâu đường ống và số lượng luồng, đồng thời lọc bỏ các cấu hình vượt quá ngân sách bộ nhớ chia sẻ. Chúng ta sử dụng decorator tự động tinh chỉnh của TileLang để biên dịch, đo kiểm, xác thực và lưu vào bộ nhớ đệm nhiều lịch trình kernel cho cùng một khối lượng công việc nhân ma trận. Sau đó, chúng ta thực thi kernel đã chọn, xác minh đầu ra của nó so với PyTorch và báo cáo độ trễ cũng như thông lượng tensor-core đạt được.
Chúng ta giới thiệu quy trình gỡ lỗi và kiểm tra nội bộ của TileLang thông qua việc in dữ liệu từ phía thiết bị (device-side printing), kiểm tra mã CUDA được tạo và trình phân tích hiệu năng (profiler) tích hợp sẵn. Chúng ta kiểm tra các điểm mốc do trình biên dịch phát ra như các toán tử tensor-core, sao chép bất đồng bộ, rào cản đồng bộ hóa và các lệnh tải ma trận. Cuối cùng, chúng ta tổ chức tất cả các phần của hướng dẫn thành một trình chạy chịu lỗi (fault-tolerant runner) giúp ghi lại trạng thái thực thi, báo cáo thông tin thời gian và in ra tài liệu tham khảo lập trình TileLang cô đọng.
Tóm lại, chúng ta đã xây dựng được sự hiểu biết thực tế về cách TileLang chuyển đổi các chương trình Python cấp tile thành các GPU kernel tối ưu mà không cần phải quản lý thủ công các chỉ số luồng, bố cục dữ liệu cấp warp, lệnh tensor-core hoặc các rào cản bộ nhớ bất đồng bộ. Chúng ta đã triển khai và xác thực các kernel bao gồm các thao tác phần tử (elementwise) bị giới hạn bởi băng thông, khối lượng công việc GEMM đòi hỏi tính toán cao, các epilogue mạng thần kinh hợp nhất, các phép khử nằm trong thanh ghi và cơ chế chú ý softmax trực tuyến. Chúng ta cũng đã xem xét cách kích thước khối, mức tiêu thụ bộ nhớ chia sẻ, độ sâu đường ống, số lượng luồng, hình dạng tile và L2 swizzling ảnh hưởng đến hiệu năng trên các kiến trúc GPU khác nhau. Cuối cùng, chúng ta đã sử dụng việc kiểm tra mã nguồn được tạo, gỡ lỗi phía thiết bị, phân tích hiệu năng và tìm kiếm lịch trình tự động để thiết lập một quy trình hoàn chỉnh cho việc phát triển, xác minh, đo kiểm và tinh chỉnh các kernel TileLang tùy chỉnh.
Xem mã nguồn đầy đủ tại đây. Ngoài ra, hãy thoải mái theo dõi chúng tôi trên Twitter và đừng quên tham gia SubReddit ML 150k+ thành viên của chúng tôi và đăng ký nhận bản tin. Khoan đã! Bạn có dùng Telegram không? Bây giờ bạn cũng có thể tham gia cùng chúng tôi trên Telegram.
Cần hợp tác với chúng tôi để quảng bá GitHub Repo, trang Hugging Face, sản phẩm mới hoặc hội thảo trực tuyến của bạn? Hãy kết nối với chúng tôi.
Sana Hassan, thực tập sinh tư vấn tại Marktechpost và là sinh viên bằng kép tại IIT Madras, có niềm đam mê áp dụng công nghệ và AI để giải quyết các thách thức trong thế giới thực. Với sự quan tâm sâu sắc đến việc giải quyết các vấn đề thực tiễn, cô mang đến một góc nhìn mới mẻ cho sự giao thoa giữa AI và các giải pháp đời sống.
Bài viết được AI dịch và tổng hợp tự động từ MarkTechPost. Liên kết bài gốc ở phía trên. AIHOT.vn luôn dẫn nguồn đầy đủ — nếu bạn thấy điểm cần chỉnh sửa, hãy gửi ý kiến tại trang phản hồi.