1 điểm bởi GN⁺ 2023-12-21 | 1 bình luận | Chia sẻ qua WhatsApp
  • mamba-minimal là một dự án triển khai Mamba một cách đơn giản và tối thiểu trong một tệp PyTorch
  • Mục tiêu là tạo ra kết quả số giống với triển khai chính thức ở forward/backward pass
  • Mã được đơn giản hóa và được cấu trúc dưới dạng có chú thích để dễ đọc
  • Không bao gồm các tối ưu hóa cốt lõi của triển khai chính thức nên không mang lại tốc độ, và cũng không bao gồm việc khởi tạo tham số phù hợp
  • Demo dùng tokenizer state-spaces/mamba-370mEleutherAI/gpt-neox-20b để chạy ví dụ hoàn thành prompt

Tổng quan dự án

  • mamba-minimal là một bản triển khai tối thiểu, đơn giản của Mamba trong một tệp PyTorch
  • Mục tiêu là thể hiện cùng hành vi với triển khai chính thức bằng mã dễ đọc hơn
  • Các đặc điểm chính:
    • Kết quả số tương đương với triển khai chính thức ở forward passbackward pass
    • Mã được đơn giản hóa
    • Triển khai dễ đọc và có chú thích

Những thứ không bao gồm

  • Tốc độ không phải là mục tiêu
    • Triển khai chính thức được tối ưu hóa mạnh
    • Các tối ưu hóa đó nằm trong những đóng góp cốt lõi của bài báo Mamba
    • Triển khai này giữ phần lớn mã ở mức đơn giản để ưu tiên khả năng đọc
  • Không bao gồm khởi tạo tham số phù hợp
    • Đây được nêu là hạng mục có thể thêm vào mà không làm giảm khả năng đọc

Ví dụ sử dụng demo

  • Có thể xem ví dụ hoàn thành prompt trong demo.ipynb
  • Ví dụ sử dụng model.MambaAutoTokenizer của Hugging Face transformers
  • Mô hình và tokenizer được dùng:
    • state-spaces/mamba-370m
    • EleutherAI/gpt-neox-20b
  • Prompt ví dụ là Mamba is the, và kết quả sinh ra bao gồm câu mô tả Mamba là rắn độc

Tài liệu tham khảo

1 bình luận

 
GN⁺ 2023-12-21
Các ý kiến trên Hacker News
  • Trước đây tôi cùng đồng nghiệp đã tạo một thư viện tách riêng phần lớn mã mô hình dùng chung; dùng nó có thể triển khai nhiều mô hình trong khoảng 100 dòng, không tính Python import và chú thích
    BERT: https://github.com/explosion/curated-transformers/blob/main/...
    Llama 1/2: https://github.com/explosion/curated-transformers/blob/main/...
    MPT: https://github.com/explosion/curated-transformers/blob/main/...
    Cũng hỗ trợ các tính năng như TorchScript JIT, PyTorch flash attention

    • Chắc chắn tôi sẽ xem qua thư viện này. Không biết bạn đã xem xformers chưa
      xformers cũng xử lý vấn đề tương tự, nhưng tập trung hơn vào việc cung cấp các mô-đun Transformer hiệu năng cao bằng Triton. Tuy nhiên, việc chỉ lấy các thành phần cụ thể của thư viện để dùng không dễ, và lỗi runtime cứ xảy ra liên tục nên tạm thời tôi gác lại. Tôi đang làm một thứ dựa trên kiến trúc BERT nên sẽ tham khảo thử
    • Tôi rất ấn tượng với thư viện này. Tôi vốn không thích lắm cách triển khai của Hugging Face, còn cái này trông như một API đẹp với mức độ trừu tượng vừa đúng
      Tôi định sẽ thử dùng trong dự án tiếp theo
  • Mã Mamba gốc có nhiều tối ưu tốc độ và các yếu tố khác nên khó hiểu ngay, còn bản triển khai này có vẻ sẽ hữu ích cho việc học
    Khi suy luận từng token một, mọi thứ trở nên đơn giản hơn rất nhiều. Tôi cũng có một bản triển khai suy luận Mamba tự làm: https://github.com/rbitr/llm.f90/tree/master/ssm

    • Fortran ư. Tôi tò mò vì sao bạn dùng Fortran
      Tôi biết nó là nền tảng của nhiều mã tính toán khoa học đã được kiểm chứng lâu năm và thường được bọc để dùng qua các thư viện như PyTorch hay Numpy, nhưng hiện nay nó không phải là ngôn ngữ phổ biến. Tôi tò mò lý do bạn chọn nó
  • Có một số điểm tôi mong được giải thích về Mamba sao cho cả người không phải nhà nghiên cứu machine learning cũng hiểu được

    1. Trực giác tổng quát của mô hình không gian trạng thái vượt ra ngoài Transformer là gì
    2. Những đổi mới tiệm tiến nào khiến Mamba thành công hơn hoặc thú vị hơn các tiền thân như S4, H3, Monarch
    3. Ngoài khả năng mở rộng dưới bậc hai theo độ dài ngữ cảnh thì nó còn có ý nghĩa gì. Ví dụ nếu không quan tâm đến độ dài ngữ cảnh trên 100k token, tôi tò mò liệu với mô hình và tập dữ liệu có kích thước tương tự, Mamba có khả năng hiệu quả hơn về chi phí tính toán khi huấn luyện hay không
    • Trí tuệ của tôi còn kém xa các tác giả bài báo, nhưng tôi vẫn cố gắng để hiểu. Tôi đã học khoa học máy tính và có trực giác cơ bản về lý thuyết điều khiển cũng như hệ rời rạc theo thời gian ở mức đại học, nhưng có lẽ để hiểu đúng bài báo này thì cần học sâu hơn nhiều về mô hình không gian trạng thái
      Trực giác cốt lõi của Mamba nằm ở việc giải một vấn đề cũ của mô hình không gian trạng thái. Mô hình không gian trạng thái giỏi nén ngữ cảnh đầu vào, nhưng trong quá trình nén đầu vào vào trạng thái ẩn, thông tin cần thiết để tận dụng ngữ cảnh hiệu quả như Transformer bị xóa mất
      Lời giải là tạo ra thứ mà bài báo gọi là cơ chế chọn lọc. Cơ chế này phụ thuộc vào đầu vào, nên mỗi khi đầu vào thay đổi, mô hình có thể điều chỉnh đầu ra ở từng bước. Để làm vậy, một số biến không gian trạng thái được biến từ bất biến theo đầu vào thành phụ thuộc vào đầu vào, và gắn thêm các lớp tuyến tính để chiếu đầu vào tại mỗi thời điểm vào các biến không gian trạng thái
      Nhưng làm cho các biến không gian trạng thái phụ thuộc vào đầu vào sẽ tạo thêm overhead tính toán. Họ giải quyết điều này bằng thuật toán nhận thức phần cứng tận dụng tối đa cấu trúc bộ nhớ GPU hiện đại, tránh tối đa việc di chuyển dữ liệu ra vào HBM
      Tri Dao là người tạo ra Flash Attention, vốn cũng là một cách dùng phần cứng hiệu quả hơn trong Transformer. Đây thực sự là lĩnh vực chuyên môn của anh ấy
    • Attention tăng theo bậc hai theo độ dài ngữ cảnh, còn mạng nơ-ron hồi quy có gating (LSTM, GRU, v.v.) là tuyến tính, và các kiến trúc mới này cũng tuyến tính. Các mạng hồi quy ban đầu dùng gating để tránh gradient bùng nổ, nhưng các cách tiếp cận mới dùng lý thuyết hệ động lực bảo đảm tính ổn định, để gating không phải giải hai vấn đề cùng lúc mà có thể tập trung vào bộ nhớ
      Mamba và Based xuất hiện ngay trước NeurIPS 2023 đã đưa vào hồi tưởng liên kết đa truy vấn (MQAR), cùng tính phụ thuộc dữ liệu của gating/chọn lọc lấy cảm hứng từ multi-head Attention. Đây là hai yếu tố then chốt còn thiếu trong Hyena và các kiến trúc không gian trạng thái trước đó; nhờ vậy các mô hình mới trở nên tốt ngang Attention ở tác vụ hồi tưởng liên kết, và trong các tác vụ không phải truy hồi thì có thể thậm chí còn nhỉnh hơn Attention một chút
      Tất nhiên, chi tiết lớn của Mamba là triển khai CUDA hiệu quả. Nếu không có nó, ý nghĩa của kiến trúc này có thể giảm đi trong những công việc mà Transformer vốn đã phù hợp
      Ngay cả khi không quá lo về độ dài ngữ cảnh, vẫn có nhiều miền mới được mở ra. Phân tích chuỗi DNA là một bài toán tuyến tính có phụ thuộc dài hạn, và cũng có thể nghĩ đến cách xem ảnh, video, thông tin nhiều chiều như một luồng token. Giống như quét pixel trên màn hình CRT ngày xưa
      Một trong những giấc mơ ban đầu của AI là một quỹ đạo học tập đơn nhất của một agent liên tục tương tác với môi trường sẽ tiến hóa liên tục, và những mô hình có độ dài ngữ cảnh vô hạn như vậy có thể giúp giấc mơ đó dễ đạt hơn
      Tuy nhiên hiện tại, các ứng dụng downstream của những mô hình này vào các tác vụ thực tế quan trọng nhìn chung vẫn chưa được kiểm chứng và tinh chỉnh nhiều bằng các ứng dụng trưởng thành dựa trên Attention. Phép so sánh với các mạng hồi quy cũ có giúp ích phần nào, nhưng trong 5 năm qua mọi người đã chuyên biệt hóa quá mức vào Attention và Transformer, nên quán tính về phía Transformer là rất lớn
    • Tôi cũng muốn biết liệu Mamba có thể được huấn luyện hiệu quả hơn về tính toán trên mô hình và tập dữ liệu có kích thước tương tự hay không
      Bài báo gốc giải thích rằng sau khi các tham số được biến đổi, mô hình có thể được tính theo hai cách: hoặc như một truy hồi tuyến tính, hoặc như một tích chập toàn cục. Thông thường, trong huấn luyện khi có thể nhìn trước toàn bộ chuỗi đầu vào, họ dùng chế độ tích chập dễ song song hóa; còn trong suy luận tự hồi quy, khi xem đầu vào từng thời điểm một, họ chuyển sang chế độ hồi quy hiệu quả
      Vì vậy huấn luyện có thể song song hóa, giống chế độ lan truyền tiến song song của RetNet. Suy luận mặc định được thực hiện ở chế độ hồi quy để có ngữ cảnh dài nhất có thể, và vì không có chunking nên khó đánh giá nó sẽ ngốn bao nhiêu RAM và VRAM trong lúc suy luận
    • Video này có vẻ đúng chính xác thứ đang tìm
      Nó vừa giải thích bài báo vừa đưa ra nhiều ngữ cảnh về việc nó nằm ở đâu trong bức tranh lớn. Nghe phần triển khai lập luận khá thú vị
      https://youtu.be/ouF-H35atOY?si=y2Ckp9MCFd7ulLL3
    • Theo tôi biết, Mamba về cơ bản là phần tiếp nối của nhánh nghiên cứu mô hình không gian trạng thái có thể gọi là tích chập dài
      Thay vì Attention bậc hai tính xem mỗi token chú ý đến mọi token khác bao nhiêu, nó bằng cách nào đó tính một kernel tích chập dài bằng độ dài đầu vào rồi áp dụng conv1d
      Theo hiểu biết hạn chế của tôi, nó hơi liên quan đến việc áp dụng FFT, nhân ma trận, rồi quay lại bằng IFFT. Tôi biết là nó chạy được nhưng chậm. Có nhiều cách tính FFT, và một trong số đó là ma trận butterfly. Có lẽ chỉ là xấp xỉ, nhưng dường như đủ tốt và rất nhanh, hiệu quả trên phần cứng hiện nay
      Độ phức tạp bậc hai nghe có vẻ tệ, nhưng trên thực tế do ràng buộc phần cứng, các thuật toán dưới bậc hai thường lại chậm hơn. Vì vậy dù kỳ vọng vào mô hình không gian trạng thái là lớn, vẫn không dễ nói rằng Llama đã hết thời. Cũng chưa biết Mamba có hoạt động tốt khi mở rộng quy mô hay không, và để biết điều đó thì thực sự phải chi hàng triệu đô la cho huấn luyện. Dù vậy tôi vẫn lạc quan
      Một mô hình thú vị khác trong họ dưới bậc hai là RWKV. Rất đáng xem qua, nhưng có lẽ đã được nói đến trong podcast rồi
      Tôi tự học và cũng chỉ từng lướt qua bài báo từ trước nên có thể sai khá nhiều. Ngoài ra Attention thường có KV cache, giúp ích rất lớn cho hiệu năng, còn tôi nghĩ Mamba thì không làm được việc đó
  • Tôi bật cười ở câu “Mamba là loài rắn độc dài nhất thế giới, với chiều dài ước tính hơn 150m”
    Dù vậy bài viết thật sự rất hay, và việc tham chiếu tới bài báo arXiv giúp những người như tôi, vốn tiêu thụ các bài kiểu này thay vì tự đọc hiểu trực tiếp bài báo, cũng có thể nhìn hé vào bên trong

    • Cái tên Mamba rất hay. Vì là [S]elective [S]tructured [S]tate [S]pace [S]equence models nên thành sSSSS, nghe giống tiếng rắn rít
    • Tôi cứ tưởng loài rắn độc dài nhất là rắn hổ mang chúa. Tìm nhanh trên Google cũng ra như vậy
      Nếu sau này phải đăng đính chính cho câu đó thì chắc sẽ thú vị
  • Tôi đoán cốt lõi của thuật toán sẽ là parallel prefix scan. Có lẽ đó mới là điểm chính của Mamba
    for i in range(l):
    x = deltaA[:, :, i] * x + deltaB_u[:, :, i]
    y = einsum(x, C[:, i, :], 'b d_in n , b n -> b d_in')
    ys.append(y)

  • Có thể là một câu hỏi ngớ ngẩn, nhưng tôi tò mò độ khó khi huấn luyện mô hình Mamba trên Hugging Face
    Mô hình lớn nhất có vẻ là 2.8B; nếu huấn luyện bằng một tập dữ liệu như The Pile thì cần bao nhiêu GPU và mất bao lâu?

    • Đây là câu hỏi hay mà tôi cũng muốn biết. Câu trả lời có vẻ là nhanh hơn đáng kể so với Transformer cùng kích thước, và kết quả cuối cùng dường như cũng sẽ đạt điểm tốt hơn Transformer trên gần như mọi benchmark
      Suy luận cũng có vẻ nhanh hơn 3~5 lần trong khi chỉ dùng một nửa RAM
  • Tôi từng cố bóc tách phiên bản CUDA chính thức nhưng thất bại ngay lần thử đầu tiên rồi rốt cuộc không động tới nữa; triển khai này trông tốt hơn nhiều

  • Lại thêm một triển khai PyTorch trong một file duy nhất, thật tuyệt vời. Tôi hy vọng hlb-CIFAR10 và các dự án liên quan trước đây, cùng những ảnh hưởng đi trước như minGPT hay DawnBench, đã góp phần dù chỉ một chút vào việc thúc đẩy định dạng một file đơn giản
    Những công việc như thế này quan trọng với nghiên cứu machine learning hiệu quả, và có thể là một trong những điều quan trọng nhất có thể làm cho lĩnh vực này lúc này
    Nghiên cứu tiến lên theo tốc độ đổi mới, đổi mới tăng tốc theo nghịch đảo của thời gian chạy thí nghiệm, và điều đó rõ ràng liên quan tới độ phức tạp Kolmogorov của mã cho mục đích nghiên cứu hay hack nhanh
    Không thể nhấn mạnh đủ rằng các công cụ như vậy quan trọng với nghiên cứu đến mức nào, và cá nhân tôi chúng đã tăng tốc quá trình khám phá tri thức ra sao. Khả năng phác thảo nhanh ý tưởng trong vài phút và lập tức nhận được kết quả có tỷ lệ tín hiệu trên nhiễu cao đã trở thành yếu tố thiết yếu để nghiên cứu tiến triển
    Tôi cho rằng chưng cất tri thức và MDL(https://en.wikipedia.org/wiki/Minimum_description_length) rất quan trọng trong việc đảo ngược những trang trí thừa thãi, mớ lộn xộn và cuộc đua chủ đề giá trị thấp quá dày đặc kiểu “không muốn bị tụt lại” mà quy trình nộp và phản biện bài báo hiện nay dường như đang khuyến khích
    Gần đây, để tránh vấn đề này và hướng tới một giải pháp mở rộng tốt hơn một chút, tôi bắt đầu phát hành mã dưới dạng “code sketch” — những gist ngắn, tự chứa, chỉ một file. Nó giảm thời gian phát triển và đưa ngay cho mọi người đoạn mã chạy được, thô ráp và chưa trau chuốt, chứa ý tưởng. Đến nay có vẻ hoạt động khá tốt và tôi muốn tiếp tục
    Tôi muốn thấy nhiều mã như thế này hơn. Nếu là các nhà nghiên cứu huấn luyện dữ liệu ở quy mô lớn, thì cả cách lan truyền thông tin cũng nên hiệu quả về dữ liệu

    • 2023 là một năm thú vị chỉ riêng việc chứng kiến nghiên cứu AI diễn ra với tốc độ phi lý. Những nền tảng như ArXiV, PyTorch, GitHub, Hugging Face và mã Python nguồn mở ngắn gọn đang tăng tốc mạnh mẽ sự phát triển của lĩnh vực mới này
      Có lẽ nhân loại chưa từng phát triển thứ gì có độ phức tạp đáng kể nhanh đến vậy
      Nơi duy nhất có tốc độ tương tự có lẽ là SpaceX, năm nay họ cũng đã phóng hai tên lửa tiên tiến nhất. Tôi tò mò năm 2024 sẽ có gì
    • Có thể có một cải thiện hiệu năng nhỏ. Ở đây x_proj không có bias, nên có vẻ có thể gộp trọng số x_proj và dt_proj lại
      Nếu có yêu cầu điều chỉnh trọng số thì có lẽ có thể làm đơn giản ở runtime, và một kernel đơn cùng bias rốt cuộc có thể nhanh hơn. Tôi không chắc
  • Tôi tự hỏi đã có thảo luận về bài báo gốc chưa. Có lẽ tôi đã bỏ lỡ, nhưng nó khá thú vị
    Tôi không hiểu rõ đoạn “do thiếu triển khai hiệu quả dẫn đến thiếu bộ nhớ hoặc yêu cầu tính toán phi thực tế, nên đã thiếu kết quả đầy đủ với độ dài ngữ cảnh 8k của các baseline RWKV và RetNet, những mô hình tuần hoàn mạnh trước đó cũng có thể được diễn giải như SSM”
    RetNet không dùng nhiều bộ nhớ, và nếu dùng triển khai lan truyền tiến theo chunk thì mức dùng VRAM bị giới hạn bởi kích thước chunk. Đây là điểm cốt lõi khi kiểm thử độ dài ngữ cảnh
    Tôi tò mò liệu đã có ai thử mô hình Mamba gốc chưa. Tốc độ huấn luyện sẽ ra sao so với RetNet ở chế độ lan truyền tiến song song?

  • Tôi luôn thích những triển khai gạn lọc các thứ phức tạp xuống chỉ còn phần cốt lõi