LLama2 Baby SFT#
项目简介#
这个项目是关于一个llama2模型的预训练和微调的, 代码仓库在:baby-llama2-chinese ↗
目前只有两步, 预训练和SFT指令微调, 是比较好的入门实验
Part1的内容是对数据集的理解, 清理数据的工程和分词的调用
硬件配置#
单卡 32G显存
内存90Gplaintext文件结构#
baby-llama2-chinese/
├── model.py # LLaMA2 模型结构与推理实现
├── pretrain.py # 预训练脚本(百度百科 / 维基语料)
├── sft.py # 有监督微调(SFT)脚本
├── eval.py # 使用微调后模型的推理 / 简单评测脚本
├── eval_pretrain.py # 预训练阶段模型的评估脚本
├── dataset.py # 预训练数据集定义与加载
├── dataset_sft.py # SFT 数据集定义与加载
├── data_process.py # 预训练语料预处理脚本
├── sft_data_process.py # SFT 数据预处理脚本
├── chatglm_tokenizer/ # ChatGLM 分词器相关文件
├── data/ # 预训练语料与中间数据
├── data_clean/ # 数据清洗、日志等辅助工具
├── sft_data/ # SFT 微调所用的指令/对话数据
├── out/ # 训练输出目录(checkpoint、日志等)
├── requirements.txt # Python 依赖列表
├── README.md # 项目说明文档plaintext数据集#
数据集分为两块, 预训练数据和微调数据
预训练数据#
我们对两个数据集分别举例看一看, 一个是百度百科, 一个是Wiki中文百科, 前一个是.bin的已经分词处理之后的语料, 后一个是.json的原始文本, 下载链接可以在原始仓库的readme里面找到
对于被分词后的.bin语料, 要用np.uint16格式读取, 读出来是Vocab ID, 即嵌入前的Token ID
# 读取为 int32 类型的 token ID 序列
data = np.fromfile('./../data/baidubaike_563w_1.bin', dtype=np.uint16)
# 使用
print(data[:100]) # 前100个token
print(f"Total tokens: {len(data)}")python[30910 32632 32357 31211 32632 32357 33662 32357 54541 32632 31201 57600
32632 54746 57897 42290 32357 31155 34699 31813 31123 36776 54727 32632
32357 54568 32883 36640 31155 32632 32357 54536 54997 57818 56194 31201
44385 31201 38146 31201 54997 54581 54902 57409 31201 54997 54722 31301
47755 32217 54997 33461 31201 44578 31201 44016 31201 54729 40964 31201
54997 55055 31201 53565 54609 31155 35284 31965 55478 55210 54642 43505
54542 34177 37746 33403 31123 54570 43647 54609 33739 32222 54536 34406
31827 31123 35400 54548 35635 54706 31123 54536 31917 40968 31123 54853
55822 33052 31123 32316]
Total tokens: 721068248plaintext我们还可以用项目里给出的tokenizer来把Token ID解码成自然语言, 注意一下路径问题
import sys
sys.path.append('./..') # 添加上级目录到路径
from chatglm_tokenizer.tokenization_chatglm import ChatGLMTokenizer
# 加载 tokenizer
tokenizer = ChatGLMTokenizer(vocab_file='./../chatglm_tokenizer/tokenizer.model')
# 取前 100 个 token 试试
token_ids = data[:100].tolist()
# 解码成文本
text = tokenizer.decode(token_ids)
print(text)python红色食品:红色食品是指食品为红色、橙红色或棕红色的食品。科学家认为,多吃些红色食品可预防感冒。
红色食品有红柿椒、西红柿、胡萝卜、红心白薯、红果(山楂)、红苹果、草莓、红枣、老南瓜、红米、柿子等。
有治疗缺铁性贫血和缓解疲劳的作用,对乳腺癌等肿瘤疾病有防治作用,给人以兴奋感,有增加食欲,光洁皮肤,增强plaintext同理去读取.json的原始文本
import json
with open('./../data/wikipedia-cn-20230720-filtered.json', 'r', encoding='utf-8') as f:
data = json.load(f)
print(data[:3] if isinstance(data, list) else data) # 打印前几条看看结构python[{'completion': '昭通机场(ZPZT)是位于中国云南昭通的民用机场,始建于1935年,1960年3月开通往返航班“昆明-昭通”,
原来属军民合用机场。1986年机场停止使用。1991年11月扩建,于1994年2月恢复通航。是西南地区「文明机场」,通航城市昆明。
机场占地1957亩,飞行区等级为4C,有一条跑道,长2720米,宽48米,可供波音737及以下机型起降。机坪面积6600平方米,停机位2个,
航站楼面积1900平方米。位于城东6公里处,民航路与金鹰大道交叉处。\n航点\n客服电话\n昭通机场客服电话:0870-2830004',
'source': 'wikipedia.zh2307'}, {'completion': '我的英雄学院:英雄新世纪\n《我的英雄学院剧场版:英雄新世纪》(仆のヒーローアカデミア THE MOVIE ヒーローズ:ライジング)
是一部于2019年12月20日上映的日本动画电影,由长崎健司执导、黑田洋介编剧,改编自日本漫画家堀越耕平创作的漫画系列《我的英雄学院》,同时也是其系列第二部电影版。
\n概要\n本作电影内容同样为作者堀越耕平监修的原创故事,并表示「这部电影版某种意义上可以说是《我的英雄学院》的结局了」。\n登场角色\n制作人员\n主题曲\n 「ハイヤーグラウンド」\n
作词:片冈健太,作曲:黑田隼之介,主唱:sumika\n跨媒体展开\n集英社亦推出该片的小说版和文库版,由小说家誉司アンリ执笔著作,小说版于2019年12月20日上市。
\n 仆のヒーローアカデミア THE MOVIE ヒーローズ:ライジング\n 仆のヒーローアカデミア\u3000THE\u3000MOVIE\u3000ヒーローズ
: ライジング\u3000ノベライズ\u3000みらい文库版', 'source': 'wikipedia.zh2307'}, {'completion': '黄大仙文化公园(Wong Tai Sin Culture Park)是香港一个公园,
位于九龙黄大仙摩士公园,门牌编号为香港黄大仙区竹园大成街8号,为一个公园中的公园,入口设于东头村道及大成街。公园原址为摩士二号公园的苗圃,后来由黄大仙区民政事务总署拨款1,370万港元,
改建为一个以中国文化为主题的公园,于2008年10月落成启用。\n设施\n* 无极广场:铺砌《易经》、八卦及太极符号\n* 诗墙:刻有书法诗篇\n* 百年古井:原为石鼓垄村的水井',
'source': 'wikipedia.zh2307'}]plaintext可以看到他的结构如下:
{'completion': '文本内容', 'source': '数据来源'}plaintext知道原始数据的结构我们才能在后面把它处理成prompt的形式
微调数据#
微调数据都是.json格式, 打印出来看一下就行
import json
# 方法1:一次性读取(适合小文件)
with open('./../sft_data/alpaca_gpt4_data_zh.json', 'r', encoding='utf-8') as f:
data = json.load(f)
print(data[:3] if isinstance(data, list) else data) # 打印前几条看看结构python[{'instruction': '保持健康的三个提示。', 'input': '', 'output': '以下是保持健康的三个提示:\n\n1. 保持身体活动。
每天做适当的身体运动,如散步、跑步或游泳,能促进心血管健康,增强肌肉力量,并有助于减少体重。\n\n2. 均衡饮食。
每天食用新鲜的蔬菜、水果、全谷物和脂肪含量低的蛋白质食物,避免高糖、高脂肪和加工食品,以保持健康的饮食习惯。\n\n3.
睡眠充足。睡眠对人体健康至关重要,成年人每天应保证 7-8 小时的睡眠。良好的睡眠有助于减轻压力,促进身体恢复,并提高注意力和记忆力。'},
{'instruction': '三原色是什么?', 'input': '', 'output': '三原色通常指的是红色、绿色和蓝色(RGB)。它们是通过加色混合原理创建色彩的三种基础颜色。
在以发光为基础的显示设备中(如电视、计算机显示器、智能手机和平板电脑显示屏), 三原色可混合产生大量色彩。其中红色和绿色可以混合生成黄色,红色和蓝色可以混合生成品红色,
蓝色和绿色可以混合生成青色。当红色、绿色和蓝色按相等比例混合时,可以产生白色或灰色。\n\n此外,在印刷和绘画中,三原色指的是以颜料为基础的红、黄和蓝颜色(RYB)。
这三种颜色用以通过减色混合原理来创建色彩。不过,三原色的具体定义并不唯一,不同的颜色系统可能会采用不同的三原色。'},
{'instruction': '描述原子的结构。', 'input': '', 'output': '原子是物质的基本单位,它由三种基本粒子组成:质子、中子和电子。质子和中子形成原子核,
位于原子中心,核外的电子围绕着原子核运动。\n\n原子结构具有层次性。原子核中,质子带正电,中子不带电(中性)。原子核非常小且致密,占据了原子总质量的绝大部分。
电子带负电,通常围绕核运动,形成若干层次,称为壳层或电子层。电子数量与质子数量相等,使原子呈电中性。\n\n电子在每个壳层中都呈规律分布,
并且不同壳层所能容纳的电子数也不同。在最里面的壳层一般只能容纳2个电子,其次一层最多可容纳8个电子,再往外的壳层可容纳的电子数逐层递增。
\n\n原子核主要受到两种相互作用力的影响:强力和电磁力。强力的作用范围非常小,主要限制在原子核内,具有极强的吸引作用,使核子(质子和中子)紧密结合在一起。
电磁力的作用范围较大,主要通过核外的电子与原子核相互作用,发挥作用。\n\n这就是原子的基本结构。原子内部结构复杂多样,不同元素的原子核中质子、中子数量不同,
核外电子排布分布也不同,形成了丰富多彩的化学世界。'}]plaintext可以看到他的结构如下:
{'instruction': '提示内容', 'input': '输入内容', 'output': '输出内容'}plaintext无input因为这不是一个续写问题, 这里需要的只是模型直接回答问题
数据清洗(Optional)#
/data_clean/clear.py里面有一大堆数据清理的函数, 用来清理原始的.json语料, 不过并不是每个函数都用到了, 我们这里详细讲一下process_baike()函数对百度百科格式数据的处理
储存格式#
处理完毕后, 会储存为.parquet格式, 等待进一步被Tokenizer处理
删除输出目录#
如果输出目录存在, 则先询问用户是否要删除
# clear.py
def delete_file(file: str)-> bool:
'''
询问删除文件
'''
if exists(file):
ans = input('delete file: {} ? Yes (y) or No (n)'.format(file))
ans = ans.lower()
if ans in ('yes', 'y'):
remove(file)
print('deleted.')
return True
return Falsepython再从process_baike()函数里面调用这个函数
# clear.py
def process_baike(response_less_word: int=15) -> None:
file_names = [
'../data/563w_baidubaike/563w_baidubaike.json',
]
save_file_name = '../data/563w_baidubaike/baike.parquet'
if exists(save_file_name):
assert delete_file(save_file_name)python去除重复的标点符号#
# clear.py
def remove_duplicate_punctuation(sentence: str) -> str:
'''
删除句子中重复的标点符号、重复的空格,同时将换行变为特殊字符'\n'
'''
# 将空格(全角空格)替换为逗号, 可能会有重复的空格,下面删除重复标点会删除
sentence = re.sub(' | ', ',', sentence)
ans = ''
n = len(sentence)
p = 0
while p < n:
ans += sentence[p]
while p + 1 < n and sentence[p] in punctuation and sentence[p + 1] in punctuation:
p += 1
p += 1
return anspython此函数把逗号变为空格, 同时删除重复的标点符号
接下来还要根据原始json数据的格式去应用这个函数
{
"title": "文章标题",
"summary": "文章简介/摘要",
"sections": [
{
"title": "章节标题1",
"content": "章节内容..."
},
{
"title": "章节标题2",
"content": "章节内容..."
}
],
"tags": ["标签1", "标签2", ...],
"url": "原始链接"
}plaintext# process_baike()函数中
def process_function(line: str) -> dict:
item = ujson.loads(line)
item_title = item['title']
item_sections = item ['sections']
for data in item_sections:
#print(item['completion'])
# 数据清洗
response = data['content'].replace('\r','')
response = remove_duplicate_punctuation(response)
# 剔除短数据
if len(response) < response_less_word:
return None
response = data['title']+data['content']
write_dict = {
"response": response,
}
return write_dictpython这个函数的作用是:
- 读取原始json数据
- 提取标题和内容(内容是一个嵌套结构, 所以需要遍历)
- 应用remove_duplicate_punctuation函数去把内容当中的每一个元素做处理
- 剔除短数据
- 返回处理后的数据
逐个处理文件#
首先需要集成之前实现的对每个文件的处理函数, 把它当作一个回调函数传入read_and_write_template_baike()函数里面去, 这样这个函数就会一次处理一个文件并且写入.parquet
# clear.py
def read_and_write_template_baike(read_file: str, write_to_file: str, call_back: object, group_cnt: int=10000) -> None:
'''
处理数据读写模板,需要提供一个回调函数call_back,
read_file: 原始数据文件
write_to_file:处理后的要保存数据文件
call_back:函数输入一个字符串,输出一个处理后的字典dict,如果输入的字符串为无效数据,请返回None
group_cnt: parquet file分割行数
如:
>>> def call_back(inputs: str) -> dict:
>>> if check(inputs) not valid:
>>> return None
...
... do something for inputs
...
>>> my_dict = {
>>> 'prompt': inputs['p'],
>>> 'response': inputs['a1'] + inputs['a2'],
>>> ...
>>> }
>>> return my_dict
'''
log.info('process file:{}'.format(read_file), save_to_file=True)
start = time.time()
raw_line_cnt = 0
keep_line_cnt = 0
with progress.open(read_file, 'r', encoding='utf-8') as f_read:
cur_rows = []
append = cur_rows.append
for line in f_read:
try:
#print(line)
raw_line_cnt += 1
write_dict = call_back(line)
if write_dict is None: continue
keep_line_cnt += 1
append(write_dict)
if len(cur_rows) >= group_cnt:
df = pd.DataFrame(cur_rows)
write_single_parquet_file(write_to_file, df)
cur_rows = []
append = cur_rows.append
except Exception as e:
# log.error('处理文件异常:{}, content:{}'.format(str(e), line))
print(line)
raise e
# end for
# 处理末尾部分
if len(cur_rows) > 0:
df = pd.DataFrame(cur_rows)
write_single_parquet_file(write_to_file, df)
cur_rows = []
end = time.time()
log.info('原始文件:{},共{}行,处理后剩余{}行,保存到文件:{}。耗时:{:.6}s'\
.format(read_file, raw_line_cnt, keep_line_cnt, write_to_file, end - start), save_to_file=True)python简而言之就是:
输入文件 (JSON/文本)
│
▼ 逐行读取
┌─────────────────┐
│ call_back() │ ← 回调函数处理每一行
└────────┬────────┘
│
┌────┴────┐
▼ ▼
None 返回 dict
(跳过) (收集保存)
│ │
└────┬────┘
▼
累计达到 group_cnt 行
│
▼
写入 Parquet 文件plaintext然后用一个for循环应用即可
# process_baike()函数
for file_name in file_names:
read_file = file_name
read_and_write_template_baike(read_file, save_file_name, process_function)python分词与编码#
本项目用的是预训练好的chatglm_tokenizer, 首先理解一下一个分词器能干什么
- tokenize(text): 分词
- encode(text): 文本转ID
- decode(ids): ID转文本
这三个是最核心的功能了, 注意分词是把一串文本分成一些子词, encode是把子词转成ID, decode是把ID转成子词
接下来这个函数就实现如何去编码这个.json语料, 还是以百度百科为例, 首先把title,summary和sections里面的内容拼接成一个str, 然后去encode这个字符串就好了
# data_process.py
def process_baidu():
BATCH_SIZE = 1000000
cnt=0
batch_cnt=0
token=0
doc_ids=[]
f1=open('./data/563w_baidubaike/563w_baidubaike.json','r',encoding='utf-8')
while True:
line = f1.readline()
if not line:
break
line=json.loads(line)
text=''
try:
text+=line['title']+':'+line['summary']
except:
pass
for per in line['sections']:
text+=per['title']+':'+per['content']+'。'
text_id=tokenizer.encode(text,add_special_tokens=False)
text_id.append(tokenizer.special_tokens['<eos>'])
if len(text_id)>5:
doc_ids+=text_id
cnt+=1
if cnt%BATCH_SIZE==0:
batch_cnt+=1
arr = np.array(doc_ids,dtype=np.uint16)
doc_ids=[]
print('cnt:',cnt,'arr_shape:',arr.shape)
with open('./data/baidubaike_563w_{}.bin'.format(batch_cnt),'wb') as f2:
f2.write(arr.tobytes())
del arr
if not doc_ids:
batch_cnt+=1
arr = np.array(doc_ids,dtype=np.uint16)
print('cnt:',cnt,'arr_shape:',arr.shape)
with open('./data/baidubaike_563w_{}.bin'.format(batch_cnt),'wb') as f:
f.write(arr.tobytes())python最后保存成.bin的格式