大模型学习(八)大模型微调之DPO训练
·
PO(直接偏好优化) 是一种无需训练奖励模型的强化学习算法,专门用于对齐大语言模型与人类偏好。
核心思想:直接比较模型对"好回答"和"坏回答"的偏好,通过优化损失函数让模型更喜欢生成人类偏好的回答。
首先定义一个创建qwen模型方法:
def create_qwen_model():
model = AutoModelForCausalLM.from_pretrained(
model_dir,
torch_dtype="auto",
device_map="auto"
)
tokenizer = AutoTokenizer.from_pretrained(model_dir)
return model,tokenizer
然后定义两个qwen模型,一个用于训练,一个用于参照:
# DPO训练的模型
model_pi,tokenizer=create_qwen_model()
# DPO参照的模型
model_ref,_=create_qwen_model()
其中学习模型是被训练优化的模型,学习生成更好的回答。参考模型是保持初始行为,防止模型"遗忘"或"跑偏"
再定义一个chat方法,跟之前的一样,只是多了model和tokenizer两个参数:
# 模型测试方法
def chat(prompt,tokenizer,model):
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": prompt},
]
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
#print(text)
model_inputs = tokenizer([text], return_tensors="pt").to(device)
generated_ids = model.generate(
model_inputs.input_ids,
max_new_tokens=512
)
generated_ids = [
output_ids[len(input_ids):] for input_ids, output_ids in zip(model_inputs.input_ids, generated_ids)
]
response = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
return response
定义DPO训练数据集:
dpo_train_data=[
{'prompt':'你是谁?','chosen':'通义千问','reject':'我是阿里云开发的超大规模语言模型,我叫通义千问。'},
{'prompt':'你是谁发明的?','chosen':'扣你jio哇','reject':'阿里巴巴'},
]
然后将DPO偏好数据转换为对话格式:
# 偏好数据集 -> 模型输入
def dpo_to_messages(dpo_pairs):
chosen_messages=[]
reject_messages=[]
for pair in dpo_pairs:
chosen_messages.append([
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": pair['prompt']},
{"role": "assistant", "content": pair['chosen']},
]
)
reject_messages.append([
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": pair['prompt']},
{"role": "assistant", "content": pair['reject']},
]
)
return chosen_messages,reject_messages
训练数据预处理与之前一样的:
# 训练数据预处理
def preprocess(tokenizer,batch_messages):
input_list=[]
target_list=[]
im_start=tokenizer('<|im_start|>').input_ids
im_end=tokenizer('<|im_end|>').input_ids
newline=tokenizer('\n').input_ids
pad=tokenizer('<|endoftext|>').input_ids
ignore=[-100]
for group in batch_messages:
input_ids=[]
target_ids=[]
for msg in group:
role=tokenizer(msg['role']).input_ids
content=tokenizer(msg['content']).input_ids
if msg['role'] in ['system','user']:
ignore_parts=role+newline+content
input_ids+=im_start+ignore_parts+im_end+newline
target_ids+=im_start+ignore*len(ignore_parts)+im_end+newline
else:
ignore_parts=role+newline
input_ids+=im_start+ignore_parts+content+im_end+newline
target_ids+=im_start+ignore*len(ignore_parts)+content+im_end+newline
input_list.append(input_ids)
target_list.append(target_ids)
# padding
max_len=max([len(ids) for ids in input_list])
for input_ids,target_ids in zip(input_list,target_list):
input_ids+=pad*(max_len-len(input_ids))
target_ids+=ignore*(max_len-len(target_ids))
batch_input_ids=torch.tensor(input_list,dtype=torch.long)
batch_target_ids=torch.tensor(target_list,dtype=torch.long)
batch_mask=batch_input_ids.ne(pad[0]).type(torch.long)
return batch_input_ids,batch_target_ids,batch_mask
模型设置为train模式:
model_pi.train()
model_ref.train()
看下结构:
Qwen2ForCausalLM(
(model): Qwen2Model(
(embed_tokens): Embedding(151936, 896)
(layers): ModuleList(
(0-23): 24 x Qwen2DecoderLayer(
(self_attn): Qwen2SdpaAttention(
(q_proj): Linear(in_features=896, out_features=896, bias=True)
(k_proj): Linear(in_features=896, out_features=128, bias=True)
(v_proj): Linear(in_features=896, out_features=128, bias=True)
(o_proj): Linear(in_features=896, out_features=896, bias=False)
(rotary_emb): Qwen2RotaryEmbedding()
)
(mlp): Qwen2MLP(
(gate_proj): Linear(in_features=896, out_features=4864, bias=False)
(up_proj): Linear(in_features=896, out_features=4864, bias=False)
(down_proj): Linear(in_features=4864, out_features=896, bias=False)
(act_fn): SiLU()
)
(input_layernorm): Qwen2RMSNorm()
(post_attention_layernorm): Qwen2RMSNorm()
)
)
(norm): Qwen2RMSNorm()
)
(lm_head): Linear(in_features=896, out_features=151936, bias=False)
)
优化器设置:
# 优化器,只训练pi模型
optimizer=torch.optim.SGD(model_pi.parameters(),lr=1e-3)
接下来就是DPO的重点,损失函数的设计:
# DPO损失计算-辅助函数
def dpo_prob_calc(target_ids,pi_logits,ref_logits):
pi_probs=torch.log_softmax(pi_logits,dim=-1) # softmax概率+log对数
ref_probs=torch.log_softmax(ref_logits,dim=-1)
ignore_mask=target_ids!=-100 # ignore token掩码
indexes=target_ids*ignore_mask # 将-100变成0,以便后面gather可以运行
pi_probs_of_target=torch.gather(pi_probs,dim=-1,index=indexes.unsqueeze(-1)).squeeze(-1) * ignore_mask # 取目标target token的概率,忽略-100 token
ref_probs_of_target=torch.gather(ref_probs,dim=-1,index=indexes.unsqueeze(-1)).squeeze(-1) * ignore_mask
pi_final_prob=pi_probs_of_target.sum(-1)/ignore_mask.sum(-1) # 求每一个样本的token prob均值
ref_final_prob=ref_probs_of_target.sum(-1)/ignore_mask.sum(-1)
return pi_final_prob,ref_final_prob
# DPO损失函数 https://github.com/huggingface/trl/blob/main/trl/trainer/dpo_trainer.py
def dpo_loss(params):
## 两个模型的chosen输出
chosen_target_ids=params['chosen_target_ids'][:,1:]
pi_chosen_logits=params['pi_chosen_logits'][:,:-1,:]
ref_chosen_logits=params['ref_chosen_logits'][:,:-1,:]
pi_chosen_prob,ref_chosen_prob=dpo_prob_calc(chosen_target_ids,pi_chosen_logits,ref_chosen_logits)
## 两个模型的reject输出
reject_target_ids=params['reject_target_ids'][:,1:]
pi_reject_logits=params['pi_reject_logits'][:,:-1,:]
ref_reject_logits=params['ref_reject_logits'][:,:-1,:]
pi_reject_prob,ref_reject_prob=dpo_prob_calc(reject_target_ids,pi_reject_logits,ref_reject_logits)
# 计算DPO Loss
pi_prob_diff=pi_chosen_prob-pi_reject_prob
ref_prob_diff=ref_chosen_prob-ref_reject_prob
beta=0.1
loss=-torch.nn.functional.logsigmoid(beta*(pi_prob_diff-ref_prob_diff))
return loss.mean()
DPO的损失函数的公式长这样:

接下来开始训练:
iterators=20
vocab=tokenizer.get_vocab()
# print(vocab)
for i in range(iterators):
# 一批模拟数据
chosen_messages,reject_messages=dpo_to_messages(dpo_train_data)
# model输入和输出
chosen_input_ids,chosen_target_ids,chosen_mask=preprocess(tokenizer,chosen_messages)
reject_input_ids,reject_target_ids,reject_mask=preprocess(tokenizer,reject_messages)
# model_pi预测
pi_chosen_logits=model_pi(input_ids=chosen_input_ids.to(device),attention_mask=chosen_mask.to(device)).logits
pi_reject_logits=model_pi(input_ids=reject_input_ids.to(device),attention_mask=reject_mask.to(device)).logits
# model_ref预测
ref_chosen_logits=model_ref(chosen_input_ids.to(device),chosen_mask.to(device)).logits
ref_reject_logits=model_ref(reject_input_ids.to(device),reject_mask.to(device)).logits
# 求DPO损失
loss=dpo_loss({
'chosen_target_ids':chosen_target_ids.to(device),
'reject_target_ids':reject_target_ids.to(device),
'pi_chosen_logits':pi_chosen_logits.to(device),
'pi_reject_logits':pi_reject_logits.to(device),
'ref_chosen_logits':ref_chosen_logits.to(device),
'ref_reject_logits':ref_reject_logits.to(device),
})
print(loss)
optimizer.zero_grad()
loss.backward()
optimizer.step()
训练的过程中只更新pi模型的权重。
model_pi.eval()
将pi模型设置为eval模式。
训练20个epoch后的效果:

感觉还没有拟合,第二个问题回答不对。继续训练30个epoch:

现在OK了,经过DPO微调之后模型原有的知识也不会发生大的变动:

更多推荐
所有评论(0)