Mistral-finetune - Fine-tune mô hình Mistral
(github.com/mistralai)mistral-finetunelà một codebase gọn nhẹ để fine-tune mô hình Mistral một cách tiết kiệm bộ nhớ và hiệu năng tốt; kho hiện đã được lưu trữ và không còn được bảo trì tích cực- Phương pháp huấn luyện dựa trên LoRA, trong đó phần lớn trọng số được cố định và chỉ huấn luyện 1~2% trọng số bổ sung dưới dạng nhiễu loạn ma trận hạng thấp
- Để đạt hiệu quả tối đa, nên dùng GPU A100 hoặc H100; mã được tối ưu cho huấn luyện multi-GPU trên một node, nhưng các mô hình nhỏ như 7B cũng có thể chạy trên một GPU
- Các mô hình được hỗ trợ gồm 7B, Mixtral 8x7B, Mixtral 8x22B, Mistral-Nemo 12B, Mistral Large v2 123B Instruct; Mistral-Nemo và Large v2 có các ràng buộc riêng về độ dài sequence và learning rate
- Dữ liệu phải theo định dạng
jsonlvà schema nghiêm ngặt; trước khi huấn luyện, việc dùngutils.validate_datađể kiểm tra định dạng và ước tính thời gian huấn luyện là bước quan trọng
Trạng thái và mục tiêu của dự án
- Kho
mistral-finetuneđang ở trạng thái Archived và không còn được bảo trì tích cực - Nếu có nhu cầu từ cộng đồng hoặc được đánh giá là có thể mang lại giá trị cho hệ sinh thái fine-tuning, trong tương lai có thể xuất hiện thư viện mới hoặc bản cập nhật lớn
- Mục tiêu là cung cấp một điểm vào đơn giản và có hướng dẫn để fine-tune các mô hình Mistral
- Codebase này khá cứng nhắc, đặc biệt về định dạng dữ liệu, và không nhắm tới việc trở thành công cụ đa dụng bao quát nhiều kiến trúc mô hình hoặc loại phần cứng
- Với cách tiếp cận tổng quát hơn, có thể tham khảo các dự án như torchtune
Phương pháp fine-tuning và khuyến nghị phần cứng
mistral-finetunedựa trên LoRA- Phần lớn trọng số mô hình được cố định
- Chỉ huấn luyện 1~2% trọng số bổ sung dưới dạng nhiễu loạn ma trận hạng thấp
- Để đạt hiệu quả tối đa, nên dùng GPU A100 hoặc H100
- Mã được tối ưu cho môi trường huấn luyện multi-GPU trên một node
- Các mô hình nhỏ như 7B có thể chạy tốt chỉ với một GPU
Cập nhật gần đây về mô hình tương thích
- Từ ngày 13 tháng 8 năm 2024, Mistral Large v2 tương thích với
mistral-finetune- Cần tải checkpoint 123B Instruct và đặt
model_id_or_paththành thư mục checkpoint đó - Do kích thước mô hình lớn, fine-tuning cần nhiều bộ nhớ hơn đáng kể
- Hiện tại cần đặt
seq_lenkhông quá 8192 - Nên dùng learning rate thấp hơn so với các mô hình khác; tài liệu gợi ý rằng trong hầu hết trường hợp
lr=1e-6sẽ hoạt động tốt
- Cần tải checkpoint 123B Instruct và đặt
- Từ ngày 19 tháng 7 năm 2024, Mistral Nemo tương thích với
mistral-finetune- Cần tải mô hình 12B Base hoặc Instruct và đặt
model_id_or_paththành thư mục checkpoint - Cần phiên bản
mistral-commonhỗ trợ Tekkenizer; cài phiên bản>=1.3.1bằngpip install --upgrade mistral-common - Do kích thước vocabulary lớn, yêu cầu bộ nhớ đỉnh của CE loss tăng lên nên hiện cần nhiều bộ nhớ hơn
- Hiện tại cần đặt
seq_lenkhông quá 16384 - Nên dùng cùng hyperparameter như 7B v3
- Cần tải mô hình 12B Base hoặc Instruct và đặt
Cài đặt và tải mô hình
- Quy trình bắt đầu gồm clone kho và cài dependency
git clone https://github.com/mistralai/mistral-finetune.gitpip install -r requirements.txt
- Khuyến nghị fine-tune các mô hình Mistral chính thức; README cung cấp các liên kết tải mô hình và checksum sau
- 7B Base:
0663b293810d7571dad25dae2f2a5806 - 7B Instruct v3:
80b71fcb6416085bcb4efad86dfb4d52 - 8x7B Base: liên kết Hugging Face
- 8x7B Instruct:
8e2d3930145dc43d3084396f49d38a3f - 8x22 Instruct:
471a02a6902706a2f1e44a693813855b - 8x22B Base:
a2fa75117174f87d1197e3a4eb50371a - 12B Instruct (Mistral-Nemo):
296fbdf911cb88e6f0be74cd04827fe7 - 12 Base (Mistral-Nemo):
c5d079ac4b55fc1ae35f51f0a3c0eb83 - 123B Instruct (Large v2):
fc602155f9e39151fba81fcaab2fa7c4
- 7B Base:
- 8x7B Base V1 và 8x7B Instruct V1 phải dùng v3 tokenizer và mở rộng kích thước vocabulary lên 32768 trước khi fine-tune
- Đường dẫn thư mục mô hình đã tải phải được chỉ định bằng đường dẫn tuyệt đối trong
model_id_or_pathcủa YAML huấn luyện
Yêu cầu về định dạng dữ liệu
- Tất cả tệp dữ liệu phải ở định dạng jsonl
- Dữ liệu tiền huấn luyện lưu văn bản thường trong khóa
"text" - Dữ liệu instruction lưu danh sách hội thoại trong khóa
"messages"- Mỗi mục chứa các khóa
"content"và"role" "role"là một trong"user","assistant","system"- Loss chỉ được tính khi
"role" == "assistant" - Có thể đặt
"weight": 0cho tin nhắn assistant để loại tin nhắn đó khỏi huấn luyện
- Mỗi mục chứa các khóa
- Dữ liệu gọi hàm cũng lưu danh sách hội thoại trong khóa
"messages"- Mỗi mục chứa khóa
"role"và"content"hoặc"tool_calls" "role"là một trong"user","assistant","system","tool"- Loss chỉ được tính khi
"role" == "assistant" "id"và"tool_call_id"trong"tool_calls"phải là chuỗi ngẫu nhiên có đúng 9 ký tự- README khuyến nghị tự động tạo các giá trị này trong script chuẩn bị dữ liệu
- Mỗi mục chứa khóa
Kiểm tra dữ liệu và workflow ví dụ
- Trước khi bắt đầu huấn luyện, cần dùng utils.validate_data để kiểm tra định dạng dữ liệu và ước tính thời gian huấn luyện
- Ví dụ instruction dùng một phần của Ultachat_200k
- Tải dữ liệu parquet bằng Pandas
- Chia 95% cho huấn luyện, 5% cho đánh giá
- Lưu thành
jsonl - Chỉ định đường dẫn trong
data.instruct_datavàdata.eval_instruct_datacủaexample/7B.yaml
- Trong quá trình kiểm tra, có thể phát hiện một số hội thoại kết thúc bằng vai trò
user- Vì chỉ huấn luyện tin nhắn assistant, tin nhắn
usercuối cùng trở thành đối tượng xử lý không cần thiết - Có thể sửa dữ liệu bằng utils.reformat_data.py
- Vì chỉ huấn luyện tin nhắn assistant, tin nhắn
- Sau khi sửa và kiểm tra lại, sẽ in ra tóm tắt như số token dữ liệu, số token huấn luyện, số epoch,
max_steps, thời gian dự kiến - Trong ví dụ README,
max_steps=500duyệt dataset khoảng 5 lần và mất khoảng 30 phút trên cụm 8xH100; cấu hình được khuyến nghị làmax_steps=300
Ví dụ fine-tuning cho gọi hàm
- Ví dụ gọi hàm dùng Glaive function calling dataset
- Dữ liệu được tải bằng Pandas, chia 95% cho huấn luyện và 5% cho đánh giá rồi lưu thành
jsonl - Dataset gốc không tuân theo định dạng gọi hàm được yêu cầu nên cần reformat
- Cần đổi
"from"thành"user" - Cần loại bỏ các ký tự
"\n"không cần thiết
- Cần đổi
- Dùng utils.reformat_data_glaive.py có thể chuyển phần lớn mẫu về đúng định dạng
- Không thể viết một script reformat hoạt động cho mọi loại dataset, nên với dataset không tuân theo định dạng yêu cầu có thể cần script reformat riêng
- Dùng
utils.validate_data --create_correctedcó thể loại bỏ các lỗi còn lại và tạo dataset.corrected
Chạy huấn luyện và ví dụ kết quả
- Sau khi kiểm tra dữ liệu, có thể bắt đầu huấn luyện
- Để huấn luyện nhanh hơn, khuyến nghị cấu hình
max_stepslà 300 run_dirphải được đặt thành thư mục thí nghiệm; có thể tùy chọn đặtwandb.projectđể dùng logging Weights & Biases- Huấn luyện được chạy bằng
torchrun, và--nproc-per-nodephải được đặt bằng số GPU khả dụng - Huấn luyện UltraChat mất khoảng 30 phút trên node 8xH100, và trọng số kết quả có thể đạt điểm MT Bench khoảng 6.3
- Huấn luyện Glaive mất khoảng 1 giờ trên node 8xH100, và tài liệu cho biết trọng số kết quả hoạt động tốt với gọi hàm
Các mục chính trong cấu hình huấn luyện
model_id_or_path: mô hình tiền huấn luyện để bắt đầu huấn luyện hoặc đường dẫn thư mục mô hình cục bộrun_dir: thư mục lưu checkpoint và metricseq_len: độ dài sequence huấn luyện; các mẫu được đóng gói theo độ dàiseq_lenđể tăng hiệu quảbatch_size: số ví dụ huấn luyện trên mỗi GPU- Kích thước batch token hiệu dụng tổng thể là
num_gpus x batch_size x seq_len
- Kích thước batch token hiệu dụng tổng thể là
max_steps: tổng số vòng lặp huấn luyện- Tổng số token được thấy trong huấn luyện là
max_steps x num_gpus x batch_size x seq_len
- Tổng số token được thấy trong huấn luyện là
optim.lr: learning rate ban đầu của optimizeroptim.weight_decay: weight decay; README khuyến nghị giữ 0.1optim.pct_start: tỷ lệ giai đoạn warm-up của PyTorchOneCycleLRlora.rank: kích thước adapter LoRA; khuyến nghị không quá 64seed: seed ngẫu nhiên để tái lập quá trình khởi tạo và xáo trộn/lấy mẫu dữ liệudata.instruct_data: đường dẫn dữ liệu instruction cho huấn luyện- Có thể chỉ định một tệp
jsonl, thư mụcjsonl, hoặc nhiều nguồn dữ liệu kèm trọng số
- Có thể chỉ định một tệp
data.data: đường dẫn dữ liệu tiền huấn luyện bổ sung tùy chọndata.eval_instruct_data: đường dẫn dữ liệu instruction đánh giá tùy chọneval_freq,no_eval,ckpt_freq: kiểm soát chu kỳ đánh giá, đánh giá trung gian và lưu checkpointsave_adapters: quyết định chỉ lưu checkpoint LoRA hay merge LoRA vào mô hình nền rồi lưu thành mô hình đầy đủsave_adapters=Falsecần đủ bộ nhớ CPU và GPU để lưu toàn bộ mô hình trong một process, thường chỉ khả thi với mô hình 7B
Suy luận và Weights & Biases
- Khuyến nghị dùng mistral-inference để suy luận với mô hình đã huấn luyện
- Có thể cài bằng
pip install mistral_inference - Khi chạy
mistral-chat, có thể dùng trọng số LoRA bằng cách chỉ định đường dẫnlora.safetensorsđã lưu trong--lora_path - Có hỗ trợ Weights and Biases để giám sát metric huấn luyện và thí nghiệm
- Cài bằng
pip install wandb - Khuyến nghị cung cấp API key qua biến môi trường
WANDB_API_KEY - Vì lý do bảo mật, API key không được đọc từ cấu hình YAML
- Training loss, evaluation loss, learning rate, v.v. được ghi và trực quan hóa trên dashboard dự án wandb
- Cài bằng
- Có thể tham khảo Weights and Biases documentation để biết cách dùng chi tiết
Mở rộng mô hình và FAQ
- Chỉ có thể fine-tune các mô hình Mistral tương thích với v3 tokenizer
- Mô hình tương thích phải có kích thước vocabulary là 32768, không phải 32000
- Các mô hình cũ có kích thước vocabulary 32000 có thể được mở rộng lên 32768 bằng
utils.extend_model_vocab - Fine-tuning mô hình MoE cho thấy độ phân tán hiệu năng lớn hơn
- Gợi ý chạy cùng một fine-tuning MoE nhiều lần với các seed khác nhau và chọn kết quả có hiệu năng tốt nhất
- Không quan sát thấy độ phân tán cao như vậy ở mô hình dense
- Có thể kiểm tra số token dùng cho huấn luyện bằng cách nhập tệp YAML huấn luyện vào utils.validate_data.py
- Nếu gặp lỗi CUDA out-of-memory, có thể giảm kích thước batch trên mỗi GPU
- Kích thước batch là
seq_len x batch_size - Gợi ý đặt
batch_sizethành 1 và giảmseq_len
- Kích thước batch là
- Thư viện được cung cấp theo Apache 2.0 License
- Không được sử dụng thư viện hoặc mô hình này theo cách xâm phạm, lạm dụng hoặc vi phạm các quyền bao gồm quyền sở hữu trí tuệ của bên thứ ba
1 bình luận
Ý kiến trên Hacker News
Khi mô hình phát triển nhanh như vậy, liệu fine-tuning còn đáng giá không? Tôi tò mò về các trường hợp sử dụng thực tế
Ví dụ, năm ngoái Bloomberg đã huấn luyện một LLM cấp GPT-3.5 bằng dữ liệu tài chính, nhưng không lâu sau đó GPT-4-8k đã vượt qua nó ở gần như mọi tác vụ tài chính
Cuối cùng chúng ta sẽ tập trung vào dữ liệu đánh giá chất lượng cao và kiến trúc cho phép dễ dàng chuyển sang mô hình mới
LLM chưa từng thấy kiểu chú thích này, LLM ngoài tiếng Anh cũng không phải ưu tiên hàng đầu của các công ty, và vì quyền riêng tư dữ liệu nên chỉ có thể dùng mô hình ưu tiên offline
Trong tình huống như vậy, fine-tuning một mô hình ngôn ngữ tổng quát là rất phù hợp
Nếu fine-tune bằng các thông điệp có cấu trúc sẵn, mô hình sẽ tự động tạo ra định dạng đó, nên có thể tiết kiệm rất nhiều token phải dùng để mô tả định dạng đầu ra trong mỗi prompt
Tuy vậy, fine-tuning lấp được khá nhiều khoảng cách giữa hai bên
Đây là một phạm vi hẹp, nhưng phần lớn lập trình cũng vậy. Nếu mục đích là làm cho LLM tổng quát nghiêng nhiều hơn về dữ liệu của mình thì fine-tuning có khả năng không liên quan lắm
Nhưng nếu bạn đang cố giải một vấn đề rất cụ thể nhưng vẫn mơ hồ, và LLM chỉ giải được một phần trong đó, fine-tuning rất có thể là lựa chọn tốt nhất
Nếu ứng dụng có nhiều hàm tùy chỉnh tương tác với công cụ, bạn có thể thích fine-tuning hơn là dùng token ngữ cảnh
Muốn làm việc này thì cần GPU nào? Tôi có bản laptop 3060 Ti, i9, RAM 16GB
Tôi không có quota AWS hay GCP và có nghe nói về Paperspace, nhưng tôi muốn nhanh chóng bắt đầu fine-tuning Mistral vì định dùng một số mô hình Mistral trong dự án khách hàng đang làm
GPU gaming có thể xả 300W nhiệt mà không vấn đề gì, nhưng nếu GPU laptop làm vậy thì sẽ chảy mất, và có khả năng bị giới hạn khoảng 100W
Tản nhiệt tỷ lệ trực tiếp với tốc độ
Hơn nữa, desktop còn có thể nâng cấp lên GPU nhanh hơn hoặc dùng nhiều GPU
Tuy nhiên, đặc biệt là cấu hình nhiều GPU sẽ ồn và tỏa nhiều nhiệt đến mức làm một căn phòng nóng lên rất nhanh
Nếu trong vài năm tới thời gian chạy GPU full-load không vượt quá 10%, cloud có khả năng rẻ hơn
Trang này tổng hợp yêu cầu phần cứng theo từng mô hình, và có thể lọc các mô hình dùng được bằng cách chọn VRAM và bộ nhớ hệ thống
Công ty chúng tôi đã fine-tune Mistral và Llama 3 bằng RTX4000 ở đó
Chỉ có 20GB RAM nên hơi hạn chế, nhưng với số token đầu vào lớn hơn, cách giảm mức lượng tử hóa đã giúp ích
Giờ họ cũng cung cấp thuê theo giờ
Công ty hiện đang dùng và đã thấy kết quả khá tốt
Rất thú vị khi xem công cụ nào sẽ trở thành tiêu chuẩn trên thực tế cho từng trường hợp sử dụng LLM phổ biến
Hệ sinh thái bị phân mảnh quá nhiều nên cảm giác như hầu hết công cụ tôi còn chưa từng nghe tới
Vài ngày trước tôi thấy Olive của Microsoft và đó là một công cụ hoàn toàn mới với tôi
Khi đã có nhiều LLM nguồn mở đạt mức “dùng được”, điều quan trọng là làm cho việc phát triển xung quanh chúng trở nên dễ dàng
Đặc biệt, những người vừa là người dùng vừa là nhà phát triển cần có khả năng tận dụng dữ liệu không công khai, chính xác hơn là dữ liệu không nằm trong quá trình tiền huấn luyện của mô hình
Kho lưu trữ nói rằng nó được tối ưu cho mô hình lớn và cần A100/H100, nhưng tôi vẫn cảm thấy việc này có thể hữu ích cho mô hình nhỏ hơn nhiều hơn là mô hình lớn
Có thể mở rộng “cứ xây rồi người ta sẽ đến” thành “cung cấp công cụ thì người ta sẽ tạo ra”
Phần trọng số rất thú vị
SFTTrainer của HuggingFace cho phép chỉ huấn luyện phần hoàn thành nếu muốn, nhưng dù điều đó trông tự nhiên với con người, LLM thường được huấn luyện tốt hơn khi dự đoán toàn bộ đầu vào
Cách này cho phép có được ưu điểm của cả hai phía
Có thể tối ưu để huấn luyện các biến thể mô hình lớn hơn bằng hai chiếc 3090 hoặc 4090 không?
Một điểm khởi đầu để xử lý vài lựa chọn là ở đây: https://huggingface.co/blog/trl-peft
Tôi có thể huấn luyện mô hình chat WhatsApp của mình như thế nào?
Bạn muốn huấn luyện mô hình bằng tin nhắn WhatsApp của mình à? Mục đích là gì? Sẽ khác nhau tùy việc bạn muốn nó viết giống bạn, hay muốn làm hỏi đáp dựa trên RAG