• Stars
    star
    121
  • Rank 293,924 (Top 6 %)
  • Language
    Python
  • License
    MIT License
  • Created over 1 year ago
  • Updated over 1 year ago

Reviews

There are no reviews yet. Be the first to send feedback to the community and the maintainers!

Repository Details

A full pipeline to finetune ChatGLM LLM with LoRA and RLHF on consumer hardware. Implementation of RLHF (Reinforcement Learning with Human Feedback) on top of the ChatGLM architecture. Basically ChatGPT but with ChatGLM

ChatGLM-LoRA-RLHF-PyTorch

a full pipeline to finetune ChatGLM LLM with LoRA and RLHF on consumer hardware


Table of Contents


Environment Setup

穷人卡:2080Ti 12G
torch==2.0.0
cuda==11.8

Todo List

  • SFT: Supervised Finetune
  • Merge Adapter into Model
  • RLHF
    • train reward model
    • tuning with RL

Run


Data Process

转化alpaca数据集为jsonl

python cover_alpaca2jsonl.py --data_path data/alpaca_data.json --save_path data/alpaca_data.jsonl

tokenization

python tokenize_dataset_rows.py --jsonl_path data/alpaca_data.jsonl --save_path data/alpaca --max_seq_length 200 --skip_overlength True

Supervised Finetune

must use latest peft version

pip uninstall peft -y
pip install git+https://github.com/huggingface/peft.git  # 最新版本 >=0.3.0.dev0
python supervised_finetune.py --dataset_path data/alpaca --lora_rank 8 --per_device_train_batch_size 1 --gradient_accumulation_steps 32 --save_steps 200 --save_total_limit 3  --learning_rate 1e-4 --fp16 --remove_unused_columns false --logging_steps 10 --output_dir output

Merge PEFT adapter into Model

pip uninstall peft -y
pip install peft==0.2.0  # 0.3.0.dev0 raise many errors
python merge_peft_adapter.py --model_name ./output 

Reward Modeling

python train_reward_model.py --model_name 'THUDM/chatglm-6b' --gradient_accumulation_steps 32 --per_device_train_batch_size 1 --train_subset 100 --eval_subset 10 --local_rank 0 --bf16 False

merge reward model into Model

python merge_peft_adapter.py --model_name ./reward_model_chatglm-6b

Notes

  1. PEFT的版本,目前从git上安装的是 0.3.0.dev0 版本,在merge_peft_adapter的时候有问题,需要切换到peft==0.2.0 (0.3.0.dev0 没有 _get_submodules()这个函数)
  2. 因为huggingface的transformer暂时不支持ChatGLM的封装接口,需要自己从ChatGLM的hub上下载代码放到本地目录 models 下面,供后续使用
  3. 同样,ChatGLM的model代码是自己的,和huggingface没合并,所以在调用加载的时候,都主要加上参数 trust_remote_code=True
  4. 训练 Reward Model 需要执行 SeqCLS 这个Task: huggingface 的 transformer 提供 "AutoModelForSequenceClassification" 这个类。但是 ChatGLM 只有 "ChatGLMForConditionalGeneration" 这个类。
  5. 自己实现 Reward model, reward_model.py,完成奖励模型的训练过程

Reference

data preprocess: cover_alpaca2jsonl.pytokenize_dataset_rows.py 来自项目 ChatGLM-Tuning

requirements 主要是按照 alpaca-lora 来配环境。


Star-History

star-history


Donation

If this project help you reduce time to develop, you can give me a cup of coffee :)

AliPay(支付宝)

ali_pay

WechatPay(微信)

wechat_pay

License

MIT © Kun

More Repositories

1

awesome_LLMs_interview_notes

LLMs interview notes and answers:该仓库主要记录大模型(LLMs)算法工程师相关的面试题和参考答案
1,126
star
2

CycleGAN-VC2

Voice Conversion by CycleGAN (语音克隆/语音转换): CycleGAN-VC2
Python
521
star
3

Vicuna-LoRA-RLHF-PyTorch

A full pipeline to finetune Vicuna LLM with LoRA and RLHF on consumer hardware. Implementation of RLHF (Reinforcement Learning with Human Feedback) on top of the Vicuna architecture. Basically ChatGPT but with Vicuna
Python
207
star
4

Recurrent-LLM

The open-source LLM implementation of paper: RecurrentGPT: Interactive Generation of (Arbitrarily) Long Text. AI 写小说,AI写作
Python
152
star
5

SecBERT

pretrained BERT model for cyber security text, learned CyberSecurity Knowledge
Python
144
star
6

CycleGAN-VC3

Voice Conversion by CycleGAN (语音克隆/语音转换):CycleGAN-VC3
Python
137
star
7

LAS_Mandarin_PyTorch

Listen, attend and spell Model and a Chinese Mandarin Pretrained model (中文-普通话 ASR模型)
Python
121
star
8

NLP4CyberSecurity

NLP model and tech for cyber security tasks
Jupyter Notebook
75
star
9

ThreatReportExtractor

Extracting Attack Behavior from Threat Reports
Python
74
star
10

Alpaca-LoRA-RLHF-PyTorch

A full pipeline to finetune Alpaca LLM with LoRA and RLHF on consumer hardware. Implementation of RLHF (Reinforcement Learning with Human Feedback) on top of the Alpaca architecture. Basically ChatGPT but with Alpaca
Python
54
star
11

nude-detect

Porn Content Pic or Video Recognization
Python
38
star
12

location_clustering

用户地理位置的聚类算法实现—基于DBSCAN和Kmeans的混合算法
Python
25
star
13

awesome_NLP-Interview-Notes

nlp_interview notes and answers: 该仓库主要记录 NLP 算法工程师相关的面试题和参考答案
18
star
14

AI-WAF

AI driven Web Application Firewall
Python
18
star
15

apk-view-tracer

Apk-view-tracer is a trigger tool for Android Dynamic Analysis and can be used in android anti-virus dynamic analysis.
Python
18
star
16

drowsiness-detection

打瞌睡检测,通过检测眼皮对眼球的遮挡程度,判定是否打瞌睡😂
Python
17
star
17

HomoglyphAttacksDetector

Detecting Homoglyph Attacks with CNN model using Computer Vision method
Jupyter Notebook
11
star
18

RepackagedAppDetector

Detect re-packaged app on Android based on fuzzy hash of instructions in dex
8
star
19

Loss-Function-In-PyTorch

Loss Function in PyTorch
Jupyter Notebook
7
star
20

WindowsStoreCrawler

crawl windows application from windows store on windows 8
C#
7
star
21

SpeakerRecognition-ResNet-GhostVLAD

Utterance-level Aggregation For Speaker Recognition In The Wild, using a "thin-ResNet" trunk architecture, and a dictionary-based NetVLAD or GhostVLAD layer to aggregate features across time, that can be trained end-to-end
Python
7
star
22

PrivacyLeakAdvancedDetection

Privacy Leak and Behavior Detect on Android based on method call graph
Java
6
star
23

audio_classification_models.pytorch

audio/voice classification in pytorch implementations
5
star
24

GANs-implementation

GAN models implementation repo
Python
4
star
25

jackaduma

personal profile
4
star
26

speaker_recognition_models.pytorch

speaker recognition / speaker verification models in pytorch implementation
4
star
27

DotNetAppGuard

Decompile &Static Analysis Dot Net App by using java
Java
4
star
28

Speech-Transformer-PyTorch

Python
4
star
29

jackaduma.github.io

CSS
4
star
30

django-cache-machine-mongoengine

Automatic caching and invalidation for Django & Mongodb. using models through the mongoengine ORM.
Python
4
star
31

py-recommender-framework

Recommender Framework implemented by python
Python
4
star
32

LangChain-OpenLLMs

Langchain-OpenLLMs with local knowledge library based on open source LLMs.
Jupyter Notebook
4
star
33

Annotated-Diffusion-Model

The Annotated Diffusion Model
Jupyter Notebook
3
star
34

SecCopilot

2
star
35

malicious-url-detection-with-ML

malicious url detection with machine learning
Python
1
star
36

awesome_AI_in_CyberSecurity_papers

awesome AI in CyberSecurity papers list
1
star
37

phishing-url-detection-with-ML

phishing url detection with machine learning
1
star
38

awesome_AI_in_Speech_papers

awesome AI in Speech papers
1
star
39

weak-password-detection-with-ML

weak password detection with machine learning
1
star