实现基于BERT的MLM和NSP预训练数据集加载类
This commit is contained in:
parent
747e381472
commit
cc01ffb5ea
|
|
@ -1,9 +1,9 @@
|
|||
## 数据集下载
|
||||
https://cims.nyu.edu/~sbowman/multinli/ <br>
|
||||
+ https://cims.nyu.edu/~sbowman/multinli/
|
||||
|
||||
## 数据格式
|
||||
|
||||
```text
|
||||
```txt
|
||||
{"annotator_labels": ["entailment", "neutral", "entailment", "neutral", "entailment"], "genre": "oup", "gold_label": "entailment", "pairID": "82890e", "promptID": "82890", "sentence1": " From Home Work to Modern Manufacture", "sentence1_binary_parse": "( From ( ( Home Work ) ( to ( Modern Manufacture ) ) ) )", "sentence1_parse": "(ROOT (PP (IN From) (NP (NP (NNP Home) (NNP Work)) (PP (TO to) (NP (NNP Modern) (NNP Manufacture))))))", "sentence2": "Modern manufacturing has changed over time.", "sentence2_binary_parse": "( ( Modern manufacturing ) ( ( has ( changed ( over time ) ) ) . ) )", "sentence2_parse": "(ROOT (S (NP (NNP Modern) (NN manufacturing)) (VP (VBZ has) (VP (VBN changed) (PP (IN over) (NP (NN time))))) (. .)))" }
|
||||
{"annotator_labels": ["entailment", "neutral", "entailment", "neutral", "entailment"], "genre": "oup", "gold_label": "entailment", "pairID": "82890e", "promptID": "82890", "sentence1": " From Home Work to Modern Manufacture", "sentence1_binary_parse": "( From ( ( Home Work ) ( to ( Modern Manufacture ) ) ) )", "sentence1_parse": "(ROOT (PP (IN From) (NP (NP (NNP Home) (NNP Work)) (PP (TO to) (NP (NNP Modern) (NNP Manufacture))))))", "sentence2": "Modern manufacturing has changed over time.", "sentence2_binary_parse": "( ( Modern manufacturing ) ( ( has ( changed ( over time ) ) ) . ) )", "sentence2_parse": "(ROOT (S (NP (NNP Modern) (NN manufacturing)) (VP (VBZ has) (VP (VBN changed) (PP (IN over) (NP (NN time))))) (. .)))"}
|
||||
```
|
||||
|
|
@ -12,7 +12,7 @@ https://cims.nyu.edu/~sbowman/multinli/ <br>
|
|||
+ 由于该数据集也可用于其它任务,因此除了我们需要的前提、假设两个句子及标签外,还有每个句子的语法解析结构等等;
|
||||
+ 下载完数据后解压,然后执行 `format.py` 脚本将原始数据按 `7:2:1` 划分成训练集、验证集和测试集。格式化后的数据形式如下所示:
|
||||
|
||||
```text
|
||||
```txt
|
||||
From Home Work to Modern Manufacture_!_Modern manufacturing has changed over time._!_1
|
||||
They were promptly executed._!_They were executed immediately upon capture._!_2
|
||||
```
|
||||
|
|
@ -0,0 +1,19 @@
|
|||
## 数据集下载
|
||||
+ https://github.com/chinese-poetry/chinese-poetry
|
||||
|
||||
|
||||
## 数据集划分
|
||||
+ 当前目录已经包含所有原始数据,运行 `format.py` 脚本可以将原始数据划分为训练集、验证集和测试集;
|
||||
+ **注意**:由于是预训练任务,模型不会直接应用于下游任务,所以没有必要单独划分测试集,因此这里保持测试集和验证集一致。
|
||||
|
||||
## 数据格式
|
||||
+ 划分后的数据格式如下:
|
||||
|
||||
```txt
|
||||
鼎湖龙远,九祭毕嘉觞。遥望白云乡。箫笳凄咽离天阙,千仗俨成行。圣神昭穆盛重光。宝室万年藏。皇心追慕思无极,孝飨奉尝。
|
||||
凤箫声断,缥缈溯丹邱。犹是忆河洲。荧煌宝册来天上,何处访仙游。葱葱郁郁瑞光浮。嘉酌侑芳羞。雕舆绣归新庙,百世与千秋。
|
||||
中兴复古,孝治日昭鸿。原庙饰瑰宫。金壁千门万,楹桷竟穹崇。亭童芝盖拥旌龙。列圣俨相从。共锡神孙千万寿,龟鼎亘衡嵩。
|
||||
```
|
||||
|
||||
+ 其中每行为一首词(一个段落),句子之间则通过句号 `。` 进行分割;
|
||||
+ 在实例化类 `LoadPretrainingDataset` 对象时,只需将 `dataset_name` 参数指定为 `songci` 即可将本数据集作为模型的预训练语料。
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
|
|
@ -0,0 +1,726 @@
|
|||
[
|
||||
{
|
||||
"author": "臧馀庆",
|
||||
"paragraphs": [
|
||||
"消息近春来,东风还又。",
|
||||
"先借椒盘劝金斗。",
|
||||
"坐间和气,压尽一番梅柳。",
|
||||
"掖庭频寓直,君恩厚。",
|
||||
"天两宫,南山齐寿。",
|
||||
"况有仙丹在公手。",
|
||||
"论功医国,合在药王之右。",
|
||||
"不妨千岁饮,长生酒。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "感皇恩"
|
||||
},
|
||||
{
|
||||
"author": "臧馀庆",
|
||||
"paragraphs": [
|
||||
"南岳有真仙,人间祥瑞。",
|
||||
"酒量诗豪世无比。",
|
||||
"晚年林下,做个清闲活计。",
|
||||
"诮如千岁鹤,巢云际。",
|
||||
"此日大家,广排筵会。",
|
||||
"酒劝千锺莫辞醉。",
|
||||
"昔时彭祖,寿年八百馀岁。",
|
||||
"十分才一分,那里暨。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "感皇恩"
|
||||
},
|
||||
{
|
||||
"author": "臧馀庆",
|
||||
"paragraphs": [
|
||||
"交广出沉香,路遥难致。",
|
||||
"何况卑人更不易。",
|
||||
"寿星香帕,我又几曾识置。",
|
||||
"有般祝寿底,忒戏。",
|
||||
"剪下一张,池州表纸。",
|
||||
"拈得轻圆更滑腻。",
|
||||
"五双纸拈,管打十个喷嚏。",
|
||||
"□□□□儿,一百岁。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "感皇恩"
|
||||
},
|
||||
{
|
||||
"author": "胡于",
|
||||
"paragraphs": [
|
||||
"去日清霜菊满丛。",
|
||||
"归来高柳絮缠空。",
|
||||
"长驱万里山收瘴,径度层波海不风。",
|
||||
"阴德遍,岭西东。",
|
||||
"天教慈母寿无穷。",
|
||||
"遥知今夕称觞处,衣彩还将衣绣同。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "胡于",
|
||||
"paragraphs": [
|
||||
"阿母蟠桃下记春。",
|
||||
"长沙星里寿星明。",
|
||||
"金花罗纸新裁诏,具叶傍行别绶经。",
|
||||
"同龙子,祝龟龄。",
|
||||
"天教二老鬓长青。",
|
||||
"明年今日称觞处,更有孙枝满谢庭。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "胡于",
|
||||
"paragraphs": [
|
||||
"楚楚吾家千里驹。",
|
||||
"老人心事正开渠。",
|
||||
"风流不减庭前玉,爱惜真如掌上珠。",
|
||||
"纡绿绶,荐方壶。",
|
||||
"老人沉醉弟兄扶。",
|
||||
"问将何物为儿寿,付与家传万卷书。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "胡于",
|
||||
"paragraphs": [
|
||||
"袅袅薰风响环。",
|
||||
"广寒仙子跨清鸾。",
|
||||
"谁教瑞世仪周间,自赋多才继小山。",
|
||||
"铃阁静,画堂闲。",
|
||||
"衮衣象服镇团栾。",
|
||||
"年年此日称觞处,留待菖蒲驻玉颜。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "申国章",
|
||||
"paragraphs": [
|
||||
"十载分符衣绣衣。",
|
||||
"裳阴处处浙东西。",
|
||||
"政成已可书银笔,词鹿仍堪付雪儿。",
|
||||
"春未半,日方迟。",
|
||||
"御沟金袅柳如丝。",
|
||||
"凤池留得梅花住,欲与先生荐寿卮。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "陈日章",
|
||||
"paragraphs": [
|
||||
"内乐清虚息万缘。",
|
||||
"逍遥真是地行仙。",
|
||||
"徙他玉带金鱼贵,听我纶巾羽扇便。",
|
||||
"倾九酝,祝长年。",
|
||||
"休辞潋滟十分舡。",
|
||||
"来年此日称觞处,定有重孙戏膝前。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "李景良",
|
||||
"paragraphs": [
|
||||
"清晓祥云绕碧天。",
|
||||
"老人星忽下南躔。",
|
||||
"庭兰共酌长生酒,持上华堂彩侍前。",
|
||||
"开绮席,舞朱颜。",
|
||||
"轻红莲叶荐金盘。",
|
||||
"沉香小院浑先暑,更有杯传数百年。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "张思济",
|
||||
"paragraphs": [
|
||||
"衣润红绡梅欲黄。",
|
||||
"几年欢意属华堂。",
|
||||
"红颜阿母逢占虺,班鬓儿童尽举觞。",
|
||||
"麟作脯,玉为浆。",
|
||||
"从今三万六千场。",
|
||||
"桑麻休说桑田变,自是经门日月长。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "鹧鸪天"
|
||||
},
|
||||
{
|
||||
"author": "李夫人",
|
||||
"paragraphs": [
|
||||
"幕天席地。",
|
||||
"瑞脑香浓笙歌沸。",
|
||||
"白衣轻。",
|
||||
"发霜髯照座明。",
|
||||
"轻簪小珥。",
|
||||
"却是人间真富贵。",
|
||||
"好着丹青。",
|
||||
"画与人间作寿星。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "减字木兰花"
|
||||
},
|
||||
{
|
||||
"author": "李夫人",
|
||||
"paragraphs": [
|
||||
"急鼓疏钟声报晓,楼上今朝,卷起重帘早。",
|
||||
"环珊珊香袅袅。",
|
||||
"尘埃不到如蓬岛。",
|
||||
"何用珠玑相映照。",
|
||||
"韵胜形清,自有天然好。",
|
||||
"莫向尊前辞醉倒。",
|
||||
"松枝鹤骨偏宜老。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "蝶恋花"
|
||||
},
|
||||
{
|
||||
"author": "李夫人",
|
||||
"paragraphs": [
|
||||
"新晴庭院暑风轻。",
|
||||
"瑞气朝来特地清。",
|
||||
"霞帔羽衣筠竹杖,萱堂真有寿星明。",
|
||||
"自来忠孝传家世,所以芝兰满户庭。",
|
||||
"簪笏成行称寿处,老来皆已鬓星星。"
|
||||
],
|
||||
"rhythmic": "瑞鹧鸪"
|
||||
},
|
||||
{
|
||||
"author": "张藻",
|
||||
"paragraphs": [
|
||||
"露零金井,尘清玉宇,双呈瑞新秋。",
|
||||
"佳气郁葱,祥烟缭绕,玉门初诞风流。",
|
||||
"宾客竞回眸。",
|
||||
"庆虎头食肉,燕颔封侯。",
|
||||
"骨相非凡,便宜谈笑上瀛洲。",
|
||||
"青衫莫欢淹留。",
|
||||
"有儿孙兰玉,不负箕裘。",
|
||||
"莲幕向来,花城今日,不妨小试良筹。",
|
||||
"名姓在金瓯。",
|
||||
"看佩珂鸣玉,促侍宸旒。",
|
||||
"直待功成名遂,归作赤松游。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "望海潮"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"枢庭喜庆生辰到。",
|
||||
"仙伯离蓬岛。",
|
||||
"鲁台云物正呈祥。",
|
||||
"线绣工夫从此、日添长。",
|
||||
"满斟绿醑深深劝。",
|
||||
"岁岁长相见。",
|
||||
"蟠桃结子几番红。",
|
||||
"笑赏清歌声调、叶黄锺。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "虞美人"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"香烟绕遍兰堂宴。",
|
||||
"香鸭珠帘卷。",
|
||||
"香风转後送韶音。",
|
||||
"香酝佳筵今日、庆佳辰。",
|
||||
"香山烧尽禽飞放。",
|
||||
"香袖佳人唱。",
|
||||
"香醪满满十分斟。",
|
||||
"香信传时延寿、保千春。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "虞美人"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"寿香烟篆金炉细。",
|
||||
"寿酒邀宾至。",
|
||||
"寿筵两畔列红妆。",
|
||||
"寿曲仙音一品、舞霓裳。",
|
||||
"寿星高挂生祥瑞。",
|
||||
"寿祝长生意。",
|
||||
"寿杯满劝庆遐龄。",
|
||||
"寿比南山松柏、永长春。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "虞美人"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"香云佳气交盘结。",
|
||||
"又庆生申节,满堂簪绂奉霞觞。",
|
||||
"瑞应台躔南极、见光芒。",
|
||||
"心田积累阴功足。",
|
||||
"来受人天福。",
|
||||
"从今几度见河清。",
|
||||
"笑傲壶中日月、镇长春。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "虞美人"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"昨霄南极星光现。",
|
||||
"今日开华宴。",
|
||||
"朱颜翠鬓占春多。",
|
||||
"且向博山香袅、卷金荷。",
|
||||
"龟游鹤舞千年寿。",
|
||||
"更酌千锺酒。",
|
||||
"满庭佳客捧殷勤。",
|
||||
"此日蒲菊新绿、瓮头春。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "虞美人"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"良辰佳景列华筵。",
|
||||
"笙歌奏管弦。",
|
||||
"位居荣显子孙贤。",
|
||||
"功名事双全。",
|
||||
"庆祝寿,拜尊前。",
|
||||
"重重福禄坚。",
|
||||
"愿如彭祖寿齐年。",
|
||||
"金杯劝寿仙。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "阮郎归"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"谢娘诗礼有家风。",
|
||||
"吹窗清昼同。",
|
||||
"笑言偕老与梁鸿。",
|
||||
"闺门喜气融。",
|
||||
"来月殿,下珠宫。",
|
||||
"人间春意浓。",
|
||||
"一尊仙酝祝芳丛。",
|
||||
"蟠桃千岁红。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "阮郎归"
|
||||
},
|
||||
{
|
||||
"author": "胡文卿",
|
||||
"paragraphs": [
|
||||
"薰风吹尽不多云。",
|
||||
"晓天如水清。",
|
||||
"哦松庭院忽闻笙。",
|
||||
"帘疏香篆明。",
|
||||
"兰玉盛,凤和鸣。",
|
||||
"家声留汉庭。",
|
||||
"狨鞍长傍九重城。",
|
||||
"年年双鬓青。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "阮郎归"
|
||||
},
|
||||
{
|
||||
"author": "杨道居",
|
||||
"paragraphs": [
|
||||
"气禀五行天与秀。",
|
||||
"瑞见枢延,况是黄锺奏。",
|
||||
"日影量来添午昼。",
|
||||
"柳梅消息年时侯。",
|
||||
"红粉吹香帘幕透。",
|
||||
"深院笙歌,劝饮金锺酒。",
|
||||
"乐事赏心千岁有。",
|
||||
"黑头人做三公有。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "蝶恋花"
|
||||
},
|
||||
{
|
||||
"author": "史佐尧",
|
||||
"paragraphs": [
|
||||
"柳垂金,梅褪玉。",
|
||||
"昴宿呈祥,符应生公族。",
|
||||
"盖世功名夸九牧。",
|
||||
"黼衮褒扬,庆阀辉南北。",
|
||||
"赐宫醪,分笃。",
|
||||
"天与长生,谩把仙椿祝。",
|
||||
"好继平阳腾茂躅。",
|
||||
"富贵千秋,饮听瑶池曲。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "苏幕遮"
|
||||
},
|
||||
{
|
||||
"author": "徐去非",
|
||||
"paragraphs": [
|
||||
"凤历书元,龟图画泰,瑞两两初开。",
|
||||
"舞风仙子,飞雪拥春来。",
|
||||
"平荡人间险秽,端为你没点尘埃。",
|
||||
"清明世,冰辉太洁,鸥鹭莫惊猜。",
|
||||
"儿孙,同劝寿,桃添西阆,兰绕南陔。",
|
||||
"看琼玉妆成,万瓦楼台。",
|
||||
"飘洒风流酥满盏,浑疑醉,月影蓬莱。",
|
||||
"从今去,诗高未老,点化尽多才。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "满庭芳"
|
||||
},
|
||||
{
|
||||
"author": "徐去非",
|
||||
"paragraphs": [
|
||||
"一种两容仪,红共白、交映南枝。",
|
||||
"紫霞仙指冰翁语,花如醉玉,香同臭雪,别样风姿。",
|
||||
"相守岁寒期。",
|
||||
"春造化、密与天知。",
|
||||
"羞将脂粉隈桃李,独先结实,还同戴胜,归宴瑶池。"
|
||||
],
|
||||
"rhythmic": "锦被堆"
|
||||
},
|
||||
{
|
||||
"author": "徐去非",
|
||||
"paragraphs": [
|
||||
"祥景飞光衮绣。",
|
||||
"流庆台,自是神仙胄。",
|
||||
"谁遣阳和放春透。",
|
||||
"化工重入丹青手。",
|
||||
"云筝锦瑟争为寿。",
|
||||
"玉带金鱼,共愿人长久。",
|
||||
"偷取蟠桃荐芳酒,更看南极星朝斗。"
|
||||
],
|
||||
"rhythmic": "卷珠帘・蝶恋花"
|
||||
},
|
||||
{
|
||||
"author": "舒大成",
|
||||
"paragraphs": [
|
||||
"祝寿筵开,华堂深映花如绣。",
|
||||
"瑞烟喷兽。",
|
||||
"帘幕香风透。",
|
||||
"一点台星,化作人间秀。",
|
||||
"韶音奏。",
|
||||
"两行红袖。",
|
||||
"争劝长生酒。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "点绛唇"
|
||||
},
|
||||
{
|
||||
"author": "贾少卿",
|
||||
"paragraphs": [
|
||||
"月寺星轺尘梦断,如今平地仙人。",
|
||||
"烟霞卷起旧精神。",
|
||||
"焚香金鸾,书奏玉麒麟。",
|
||||
"闻道枫宸求侍从,看庆命重新。",
|
||||
"且将风月剩身。",
|
||||
"樽中长有酒,花下不辜春。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "临江仙"
|
||||
},
|
||||
{
|
||||
"author": "陆汉广",
|
||||
"paragraphs": [
|
||||
"绿莺庭院燕莺啼。",
|
||||
"绣帘垂。",
|
||||
"瑞烟霏。",
|
||||
"一片笙箫,声过彩云低。",
|
||||
"疑是蕊宫仙子降,翻玉袖,舞瑶姬。",
|
||||
"冰姿玉质自清奇。",
|
||||
"看孙枝。",
|
||||
"列班衣。",
|
||||
"画鼓新歌,喜映两疏眉。",
|
||||
"袖里蟠桃花露湿,应不惜,醉金卮。"
|
||||
],
|
||||
"rhythmic": "江城子"
|
||||
},
|
||||
{
|
||||
"author": "王阜民",
|
||||
"paragraphs": [
|
||||
"家法从来师静治,赵张高掩前踪。",
|
||||
"清才八斗继宗风。",
|
||||
"年年正月尾,桃李满城中。",
|
||||
"已把长江成九酝,请将太白浮公。",
|
||||
"更移春槛向房栊。",
|
||||
"有花虽解语,莫负锦薰笼。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "临江仙"
|
||||
},
|
||||
{
|
||||
"author": "霍安人",
|
||||
"paragraphs": [
|
||||
"十月小春天,梅飘香细。",
|
||||
"九叶尧已呈瑞。",
|
||||
"寿阳仙子,暂降羽衣环。",
|
||||
"林间风味,别人难比。",
|
||||
"齐眉共庆,劝声鼎沸。",
|
||||
"有子知书继家世。",
|
||||
"兼金五福,行被君恩宠贲。",
|
||||
"愿祈龟鹤算,千千岁。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "感皇恩"
|
||||
},
|
||||
{
|
||||
"author": "霍安人",
|
||||
"paragraphs": [
|
||||
"正朱明时侯,院宇清和,庆逢佳节。",
|
||||
"梦应熊罴,尧翻三叶,罗绮如云,寿杯争劝,竞起歌新阕。",
|
||||
"瑞气氤氲,祥云缭绕,玉炉频。",
|
||||
"溪室封功,几多勋业,首冠今朝,一时英杰。",
|
||||
"得配侯门,岂不惭疏拙。",
|
||||
"彩凤和鸣,早膺荣擢。",
|
||||
"增盛斑衣列。",
|
||||
"福禄无穷,年过卫武,辉光阀阅。"
|
||||
],
|
||||
"rhythmic": "醉蓬莱"
|
||||
},
|
||||
{
|
||||
"author": "霍安人",
|
||||
"paragraphs": [
|
||||
"桐叶霜乾,芦花风软,晓来一色新秋。",
|
||||
"碧光无际,良夜月明楼。",
|
||||
"瑞应长庚入梦,锺奇秀、特产贤侯。",
|
||||
"堪夸处,雄姿英发,连箭射双雕。",
|
||||
"回头。",
|
||||
"思往事,皂囊三进,豪气冲牛。",
|
||||
"记前回凤诏,空下南洲。",
|
||||
"冷笑浮云坠甑,鲈鱼美、归老扁舟。",
|
||||
"祝君寿,青山不尽,绿水自悠悠。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "满庭芳"
|
||||
},
|
||||
{
|
||||
"author": "希叟",
|
||||
"paragraphs": [
|
||||
"燕堂秋未老。",
|
||||
"正木犀香散,芙蓉红小。",
|
||||
"门阑瑞烟绕。",
|
||||
"望银河月暗,寿星偏照。",
|
||||
"碧霞道要。",
|
||||
"有真人、亲传最妙。",
|
||||
"况桃源旧约,重寻鬓发,胜如年少。",
|
||||
"缥缈。",
|
||||
"六铢衣降,九转丹成,五云齐到。",
|
||||
"十洲三岛。",
|
||||
"神仙路,终须到。",
|
||||
"对芝兰玉树,宝杯交劝,何惜玉山醉倒。",
|
||||
"看乘鸾跨鹤,归来洞天未晓。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "瑞鹤仙"
|
||||
},
|
||||
{
|
||||
"author": "沈元实",
|
||||
"paragraphs": [
|
||||
"金关五云里,玉座太微间。",
|
||||
"凌虚新就燕间,宣唤侍臣班。",
|
||||
"丹坐移前席,禁漏声传高阁,喜气满龙颜。",
|
||||
"天语眷畴昔,政路稳跻攀。",
|
||||
"酒如渑,香袅穗,寿南山。",
|
||||
"橙黄桔绿,樽前辉映菊花团。",
|
||||
"清晓凉风凝露,晴昼秋光满院,岁岁奉君。",
|
||||
"待看云台画,荣观侈人寰。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "水调歌头"
|
||||
},
|
||||
{
|
||||
"author": "游子蒙",
|
||||
"paragraphs": [
|
||||
"春玉苍山,屏星暖、佳辰难得。",
|
||||
"看柳眼梅金,全似海沂春色。",
|
||||
"和气欢声俄尔许,祥烟瑞霭今何夕。",
|
||||
"是风流、儒雅黑头公,悬弧日。",
|
||||
"福不尽,贵无敌。",
|
||||
"愿岁岁,见华席。",
|
||||
"捧霞觞称寿,寿如南极。",
|
||||
"见说蟠桃花正发,柔风暖日瑶池碧。",
|
||||
"待他年、结子欲成时,留君摘。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "满江红"
|
||||
},
|
||||
{
|
||||
"author": "游子蒙",
|
||||
"paragraphs": [
|
||||
"雪坞霜林,一夜报、春归消息。",
|
||||
"看是处、春回柳眼,粉匀梅额。",
|
||||
"和气已欣回北陆,寿星更喜明南极。",
|
||||
"问四井、节物奉谁欢,辽东客。",
|
||||
"熊梦旦,非常日。",
|
||||
"珠履闹,金钗密。",
|
||||
"指双溪千顷,共斟琼液。",
|
||||
"不用殷勤千岁祝,姓名已上神仙籍。",
|
||||
"但时从、王母借蟠桃,躬亲摘。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "满江红"
|
||||
},
|
||||
{
|
||||
"author": "贾逋",
|
||||
"paragraphs": [
|
||||
"薰梅染柳。",
|
||||
"借得东君手。",
|
||||
"柳色梅香到樽前,摅写才华八斗。",
|
||||
"当年占梦佳辰。",
|
||||
"今年乐事尤新。",
|
||||
"日侍玉皇香案,钧天日月常春。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "清平乐"
|
||||
},
|
||||
{
|
||||
"author": "吴文若",
|
||||
"paragraphs": [
|
||||
"玉宇生凉秋恰半。",
|
||||
"月到今霄,分外清光满。",
|
||||
"兔魄呈祥冰烂。",
|
||||
"广寒仙子生华旦。",
|
||||
"聪慧风流天与擅,淑质冰婆,本是飞琼伴。",
|
||||
"□领彩衣椿祝劝。",
|
||||
"蟠桃待熟瑶池宴。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "蝶恋花"
|
||||
},
|
||||
{
|
||||
"author": "去非",
|
||||
"paragraphs": [
|
||||
"龙角辉春,蛾春惊晓,梦阑金翠屏开。",
|
||||
"异芬薰室,风送蕊仙来。",
|
||||
"玉女擎香沐浴,人间世、洗彻凡埃。",
|
||||
"梅开後,留花酝染,清味俗难猜。",
|
||||
"东君,尤雅爱,传香芳畹,香发庭陔。",
|
||||
"宁馨满尊前,喜奏瑶台。",
|
||||
"便好纽为佩王,瀛洲路、同赏蓬莱。",
|
||||
"蟠桃宴,从今曼倩,三骋奇材。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "满庭芳"
|
||||
},
|
||||
{
|
||||
"author": "潘熊飞",
|
||||
"paragraphs": [
|
||||
"十日後重阳。",
|
||||
"甘菊阶前满意黄。",
|
||||
"生日无钱留贺客,何妨。",
|
||||
"尚有儿曹理寿觞。",
|
||||
"双鬓已沧浪。",
|
||||
"休问金门与玉堂。",
|
||||
"二仲相期三径在,徜徉。",
|
||||
"何用功成似子房。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "南乡子"
|
||||
},
|
||||
{
|
||||
"author": "黄庭佐",
|
||||
"paragraphs": [
|
||||
"露着桂枝晓,霜护菊篱秋。",
|
||||
"无尘玉宇南极,一点瑞光浮。",
|
||||
"崧岳精英瑞世,河洛图书寓直,琳馆奉宸游。",
|
||||
"独袖功名手,谁与复神州。",
|
||||
"带垂金,头尚黑,紫绮裘。",
|
||||
"年年今日,灵寿扶出富民侯。",
|
||||
"应为娉婷一笑,烂醉葡萄新熟,明月满西楼。",
|
||||
"西北雨峰峙,端与寿山侔。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "水调歌头"
|
||||
},
|
||||
{
|
||||
"author": "赵□",
|
||||
"paragraphs": [
|
||||
"寸心千里。"
|
||||
],
|
||||
"rhythmic": "失调名"
|
||||
},
|
||||
{
|
||||
"author": "吴氏3",
|
||||
"paragraphs": [
|
||||
"剪新幡儿,斜插真珠髻。"
|
||||
],
|
||||
"rhythmic": "失调名"
|
||||
},
|
||||
{
|
||||
"author": "吴氏3",
|
||||
"paragraphs": [
|
||||
"楼台里,春风淡荡。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "南乡子"
|
||||
},
|
||||
{
|
||||
"author": "吴氏3",
|
||||
"paragraphs": [
|
||||
"乍卷珠帘新燕入。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "南乡子"
|
||||
},
|
||||
{
|
||||
"author": "吴氏3",
|
||||
"paragraphs": [
|
||||
"几声天外归鸿。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "多丽"
|
||||
},
|
||||
{
|
||||
"author": "吴氏3",
|
||||
"paragraphs": [
|
||||
"一声初报晓。",
|
||||
" >> ",
|
||||
"词牌介绍"
|
||||
],
|
||||
"rhythmic": "渔家傲"
|
||||
}
|
||||
]
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
|
|
@ -0,0 +1,41 @@
|
|||
import json
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def format():
|
||||
"""划分原始的ci.song.xxx.json数据集,保存为训练、验证和测试集"""
|
||||
|
||||
def read_file(pathname=None):
|
||||
"""读取每个json文件中的1000首词"""
|
||||
|
||||
paragraphs = []
|
||||
with open(pathname, encoding="utf-8") as f:
|
||||
json_data = json.loads(f.read())
|
||||
for item in json_data:
|
||||
para = item["paragraphs"] # 取出一首词(一个段落)的列表
|
||||
# 去除"词牌介绍"和" >> "两句
|
||||
if para[-1] == "词牌介绍":
|
||||
para = para[:-2]
|
||||
if len(para) < 2: # 舍弃小于两句的词
|
||||
continue
|
||||
paragraphs.append(para)
|
||||
|
||||
return paragraphs
|
||||
|
||||
def make_data(pathname, start, end):
|
||||
with open(pathname, "w", encoding="utf-8") as f:
|
||||
for i in tqdm(
|
||||
range(start, end, 1000), ncols=80, desc=f"## 正在制作数据集{pathname}"
|
||||
):
|
||||
json_file = f"ci.song.{i}.json"
|
||||
paragraphs = read_file(json_file)
|
||||
for para in paragraphs:
|
||||
f.write("".join(para) + "\n")
|
||||
|
||||
make_data("songci.train.txt", 0, 19001) # 20 * 1000首
|
||||
make_data("songci.valid.txt", 20000, 21001) # 2 * 1000首
|
||||
make_data("songci.test.txt", 20000, 21001) # 2 * 1000首
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
format()
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
|
@ -3,14 +3,14 @@ https://github.com/aceimnorstuvwxz/toutiao-text-classfication-dataset
|
|||
|
||||
## 数据格式
|
||||
|
||||
```text
|
||||
```txt
|
||||
6552431613437805063_!_102_!_news_entertainment_!_谢娜为李浩菲澄清网络谣言,之后她的两个行为给自己加分_!_佟丽娅,网络谣言,快乐大本营,李浩菲,谢娜,观众们
|
||||
```
|
||||
每行一条数据,以 `_!_` 分割字段,从前往后分别是:新闻ID,类别代码(见下文),类别名称(见下文),新闻标题文本,新闻关键词
|
||||
|
||||
## 类别与名称
|
||||
|
||||
```text
|
||||
```txt
|
||||
100 民生 故事 news_story
|
||||
101 文化 文化 news_culture
|
||||
102 娱乐 娱乐 news_entertainment
|
||||
|
|
@ -34,7 +34,7 @@ https://github.com/aceimnorstuvwxz/toutiao-text-classfication-dataset
|
|||
原始数据集下载完成后,运行当前文件夹中的 `format.py` 脚本文件即可将原始数据按照 `7:2:1` 的比例划分成规整的训练集 `toutiao_train.txt`、验证集 `toutiao_val.txt` 和 测试集 `test.txt`
|
||||
|
||||
处理完成后的数据格式如下:
|
||||
```text
|
||||
```txt
|
||||
轻松一刻:带你看全球最噩梦监狱,每天进几百人,审讯时已过几年_!_11
|
||||
千万不要乱申请网贷,否则后果很严重_!_4
|
||||
10年前的今年,纪念5.12汶川大地震10周年_!_11
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
from .BERT.config import BertConfig
|
||||
|
|
@ -0,0 +1,90 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.append(os.getcwd())
|
||||
|
||||
import torch
|
||||
import logging
|
||||
from utils import logger_init
|
||||
from model import BertConfig
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
|
||||
class ModelConfig:
|
||||
"""基于BERT的MLM、NSP预训练模型的配置类"""
|
||||
|
||||
def __init__(self):
|
||||
self.project_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
# # ========== wikitext2数据集相关配置 ==========
|
||||
# self.dataset_dir = os.path.join(
|
||||
# self.project_dir, "data", "mlm_nsp_pretraining", "wikitext2"
|
||||
# )
|
||||
# self.pretrained_model_dir = os.path.join(
|
||||
# self.project_dir, "archive", "bert_base_uncased_english"
|
||||
# )
|
||||
# self.train_filepath = os.path.join(self.dataset_dir, "wiki.train.tokens")
|
||||
# self.val_filepath = os.path.join(self.dataset_dir, "wiki.valid.tokens")
|
||||
# self.test_filepath = os.path.join(self.dataset_dir, "wiki.test.tokens")
|
||||
# self.dataset_name = "wikitext2"
|
||||
# self.split_sep = " . "
|
||||
|
||||
# ========== songci数据集相关配置 ==========
|
||||
self.dataset_dir = os.path.join(
|
||||
self.project_dir, "data", "pretraining", "songci"
|
||||
)
|
||||
self.pretrained_model_dir = os.path.join(
|
||||
self.project_dir, "archive", "bert_base_chinese"
|
||||
)
|
||||
self.train_filepath = os.path.join(self.dataset_dir, "songci.train.txt")
|
||||
self.val_filepath = os.path.join(self.dataset_dir, "songci.valid.txt")
|
||||
self.test_filepath = os.path.join(self.dataset_dir, "songci.test.txt")
|
||||
self.dataset_name = "songci"
|
||||
self.split_sep = "。"
|
||||
|
||||
self.vocab_path = os.path.join(self.pretrained_model_dir, "vocab.txt")
|
||||
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
self.model_save_dir = os.path.join(self.project_dir, "cache")
|
||||
if not os.path.exists(self.model_save_dir):
|
||||
os.makedirs(self.model_save_dir)
|
||||
self.model_save_path = os.path.join(
|
||||
self.model_save_dir, f"{self.dataset_name}_pretraining_model.bin"
|
||||
)
|
||||
|
||||
self.log_level = logging.INFO
|
||||
self.log_save_dir = os.path.join(self.project_dir, "logs")
|
||||
logger_init(
|
||||
log_filename=self.dataset_name + "_pretraining",
|
||||
log_level=self.log_level,
|
||||
log_dir=self.log_save_dir,
|
||||
)
|
||||
|
||||
self.epochs = 200
|
||||
self.batch_size = 32
|
||||
self.is_sample_shuffle = True
|
||||
self.max_sen_len = None # 填充模式
|
||||
self.eval_per_epoch = 1 # 验证模型的epoch数
|
||||
self.pad_index = 0
|
||||
self.random_state = 2023
|
||||
self.learning_rate = 4e-5
|
||||
self.weight_decay = 0.1
|
||||
# False表示使用自定义的多头注意力模块;True表示使用torch框架中实现的
|
||||
self.use_torch_multi_head = False
|
||||
|
||||
self.masked_rate = 0.15 # 掩码比例
|
||||
self.masked_token_rate = 0.8 # 用[MASK]token替换的比例
|
||||
self.masked_token_unchanged_rate = 0.5 # 保持不变的token比例
|
||||
self.use_embedding_weight = True
|
||||
self.writer = SummaryWriter(f"runs/{self.dataset_name}_pretraining")
|
||||
|
||||
# 导入BERT模型部分配置
|
||||
bert_config_path = os.path.join(self.pretrained_model_dir, "config.json")
|
||||
bert_config = BertConfig.from_json_file(bert_config_path)
|
||||
for key, value in bert_config.__dict__.items():
|
||||
self.__dict__[key] = value
|
||||
|
||||
# 将当前配置打印到日志文件中
|
||||
logging.info("=" * 20)
|
||||
logging.info("### 将当前配置打印到日志文件中")
|
||||
for key, value in self.__dict__.items():
|
||||
logging.info(f"### {key} = {value}")
|
||||
|
|
@ -0,0 +1,64 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.append(os.getcwd())
|
||||
|
||||
from tasks.pretraining import ModelConfig
|
||||
from utils import LoadPretrainingDataset
|
||||
from transformers import BertTokenizer
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
model_config = ModelConfig()
|
||||
dataset = LoadPretrainingDataset(
|
||||
vocab_path=model_config.vocab_path,
|
||||
tokenizer=BertTokenizer.from_pretrained(
|
||||
model_config.pretrained_model_dir
|
||||
).tokenize,
|
||||
batch_size=model_config.batch_size,
|
||||
max_sen_len=model_config.max_sen_len,
|
||||
max_position_embeddings=model_config.max_position_embeddings,
|
||||
split_sep=model_config.split_sep,
|
||||
pad_index=model_config.pad_index,
|
||||
is_sample_shuffle=model_config.is_sample_shuffle,
|
||||
dataset_name=model_config.dataset_name,
|
||||
masked_rate=model_config.masked_rate,
|
||||
masked_token_rate=model_config.masked_token_rate,
|
||||
masked_token_unchanged_rate=model_config.masked_token_unchanged_rate,
|
||||
random_state=model_config.random_state,
|
||||
)
|
||||
|
||||
train_loader, val_loader, test_loader = dataset.data_loader(
|
||||
train_filepath=model_config.train_filepath,
|
||||
val_filepath=model_config.val_filepath,
|
||||
test_filepath=model_config.test_filepath,
|
||||
)
|
||||
# # 仅生成测试集
|
||||
# test_loader = dataset.data_loader(
|
||||
# test_filepath=model_config.test_filepath, only_test=True
|
||||
# )
|
||||
|
||||
for id_seqs, segs, b_mask, mlm_labels, nsp_labels in test_loader:
|
||||
print(f"token id seqs shape: #{id_seqs.shape}") # [src_len, batch_size]
|
||||
print(f"token type id seqs shape: #{segs.shape}") # [src_len, batch_size]
|
||||
print(f"padding mask shape: #{b_mask.shape}") # [batch_size, src_len]
|
||||
print(f"MLM task labels shape: #{mlm_labels.shape}") # [src_len, batch_size]
|
||||
print(f"NSP task labels shape: #{nsp_labels.shape}") # [batch_size]
|
||||
|
||||
id_seq = id_seqs.transpose(0, 1)[0]
|
||||
mlm_label = mlm_labels.transpose(0, 1)[0]
|
||||
token = " ".join([dataset.vocab.itos[id] for id in id_seq])
|
||||
label = " ".join([dataset.vocab.itos[id] for id in mlm_label])
|
||||
print(f"Mask后的token序列:{token}")
|
||||
print(f"对应的MLM任务标签:{label}")
|
||||
|
||||
break
|
||||
|
||||
sentences = ["十年生死两茫茫。不思量。自难忘。千里孤坟,无处话凄凉。", "红酥手。黄藤酒。满园春色宫墙柳。"]
|
||||
b_id_seq, b_mask_pos, b_padding_mask = dataset.get_inference_samples(
|
||||
sentences, masked=False
|
||||
)
|
||||
print("=" * 10, "推理时示例样本", "=" * 10)
|
||||
print(f"token id seqs: {b_id_seq.transpose(0, 1)}")
|
||||
print(f"MLM task labels: {b_mask_pos}")
|
||||
print(f"padding mask: {b_padding_mask}")
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
from .log_helper import logger_init
|
||||
from .create_pretraining_data import LoadPretrainingDataset
|
||||
|
|
@ -0,0 +1,457 @@
|
|||
import os
|
||||
import random
|
||||
import logging
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
from .data_helpers import Vocab
|
||||
from .data_helpers import pad_sequence
|
||||
from .data_helpers import process_cache
|
||||
|
||||
|
||||
def format_wikitext2(filepath=None, sep=" . "):
|
||||
"""
|
||||
格式化原始的wikitext2数据集
|
||||
:return: 返回一个二维list,外层list元素为一个文本段落;内层list元素为一个段落中的句子 `[[para1_sen1, para1_sen2, ...], [para2_sen1, para2_sen2, ...], ...]`
|
||||
"""
|
||||
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
lines = f.readlines() # 读取所有行,每一行为一个文本段落
|
||||
|
||||
paragraphs = []
|
||||
for line in tqdm(lines, ncols=80, desc="## 正在读取wikitext2原始数据"):
|
||||
# 将段落转为小写,并按分隔符分为句子
|
||||
sentences = line.lower().split(sep)
|
||||
|
||||
# 给每个句子加上分隔符,并且去除最后的空句
|
||||
tmp_sens = []
|
||||
for sen in sentences:
|
||||
sen = sen.strip()
|
||||
if len(sen) == 0:
|
||||
continue
|
||||
sen += sep
|
||||
tmp_sens.append(sen)
|
||||
|
||||
# 若段落内少于两条句子,则舍弃,因为NSP任务的输入需要一对句子
|
||||
if len(tmp_sens) < 2:
|
||||
continue
|
||||
|
||||
paragraphs.append(tmp_sens)
|
||||
|
||||
random.shuffle(paragraphs) # 将所有段落打乱
|
||||
return paragraphs
|
||||
|
||||
|
||||
def format_songci(filepath=None, sep="。"):
|
||||
"""格式化原始的宋词数据集"""
|
||||
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
lines = f.readlines() # 一次读取所有行,每一行为一首词
|
||||
|
||||
paragraphs = []
|
||||
for line in tqdm(lines, ncols=80, desc="## 正在读取宋词原始数据"):
|
||||
# 去除有乱码字符的段落
|
||||
if "□" in line or "……" in line:
|
||||
continue
|
||||
|
||||
sentences = line.split(sep)
|
||||
|
||||
# 给每个句子加上分隔符,并且去除最后的空句
|
||||
tmp_sens = []
|
||||
for sen in sentences:
|
||||
sen = sen.strip()
|
||||
if len(sen) == 0:
|
||||
continue
|
||||
sen += sep
|
||||
tmp_sens.append(sen)
|
||||
|
||||
# 去除少于两个句子的段落
|
||||
if len(tmp_sens) < 2:
|
||||
continue
|
||||
|
||||
paragraphs.append(tmp_sens)
|
||||
|
||||
random.shuffle(paragraphs) # 将所有段落打乱
|
||||
return paragraphs
|
||||
|
||||
|
||||
def format_custom(filepath=None, sep=None):
|
||||
"""格式化自定义的数据集"""
|
||||
|
||||
raise NotImplementedError(
|
||||
"本函数未实现,请参照 `format_wikitext2()` 或 `format_songci()` 函数返回格式进行实现"
|
||||
)
|
||||
|
||||
|
||||
class LoadPretrainingDataset(object):
|
||||
"""加载预训练数据集"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_path="./vocab.txt",
|
||||
tokenizer=None,
|
||||
batch_size=32,
|
||||
max_sen_len=None,
|
||||
max_position_embeddings=512,
|
||||
split_sep="。",
|
||||
pad_index=0,
|
||||
is_sample_shuffle=True,
|
||||
dataset_name="wikitext2",
|
||||
masked_rate=0.15,
|
||||
masked_token_rate=0.8,
|
||||
masked_token_unchanged_rate=0.5,
|
||||
random_state=2023,
|
||||
):
|
||||
self.vocab = Vocab(vocab_path)
|
||||
self.tokenizer = tokenizer
|
||||
self.batch_size = batch_size
|
||||
|
||||
# min(max_sen_len, max_position_embeddings)决定了输入序列长度
|
||||
if isinstance(max_sen_len, int) and max_sen_len > max_position_embeddings:
|
||||
max_sen_len = max_position_embeddings
|
||||
self.max_sen_len = max_sen_len
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
|
||||
self.split_sep = split_sep
|
||||
self.PAD_IDX = pad_index
|
||||
self.CLS_IDX = self.vocab["[CLS]"]
|
||||
self.SEP_IDX = self.vocab["[SEP]"]
|
||||
self.MASK_IDX = self.vocab["[MASK]"]
|
||||
self.is_sample_shuffle = is_sample_shuffle
|
||||
|
||||
self.dataset_name = dataset_name
|
||||
self.masked_rate = masked_rate
|
||||
self.masked_token_rate = masked_token_rate
|
||||
self.masked_token_unchanged_rate = masked_token_unchanged_rate
|
||||
self.random_state = random_state
|
||||
random.seed(random_state) # 设置随机状态,用于复现结果
|
||||
|
||||
def format_data(self, filepath):
|
||||
"""
|
||||
将原始数据集格式化为标准形式
|
||||
:return: `[[para1_sen1, para1_sen2, ...], [para2_sen1, para2_sen2, ...], ...]`
|
||||
"""
|
||||
|
||||
# 依据数据集名称调用对应的格式化函数,注意:格式化函数返回格式需要保持一致
|
||||
# wikitext2数据集
|
||||
if self.dataset_name == "wikitext2":
|
||||
return format_wikitext2(filepath, self.split_sep)
|
||||
# 宋词数据集
|
||||
elif self.dataset_name == "songci":
|
||||
return format_songci(filepath, self.split_sep)
|
||||
# 其他自定义数据集
|
||||
elif self.dataset_name == "custom":
|
||||
return format_custom(filepath)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"数据集 {self.dataset_name} 不存在对应的格式化函数,"
|
||||
f"请参考函数 `format_wikitext2()` 实现对应的格式化函数!"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_next_sentence_sample(sentence, next_sentence, paragraphs):
|
||||
"""由给定的连续两个序列和所有文本段落,生成一个句子对样本"""
|
||||
|
||||
# 正负样本数量相同
|
||||
# 传入的两个句子构成正样本,标签为True
|
||||
if random.random() < 0.5:
|
||||
is_next = True
|
||||
# 构造负样本,标签为False
|
||||
else:
|
||||
# 先随机选中一个段落,再随机选中一个句子
|
||||
new_next_sentence = next_sentence
|
||||
while next_sentence == new_next_sentence: # 避免随机选中的下一个句子和传入的下一个句子相同
|
||||
new_next_sentence = random.choice(random.choice(paragraphs))
|
||||
next_sentence = new_next_sentence
|
||||
is_next = False
|
||||
|
||||
return sentence, next_sentence, is_next
|
||||
|
||||
def masking_tokens(self, token_ids, candidate_mask_positions, num_mask_ids):
|
||||
"""依据需要mask的tokens数量和候选mask位置,对token_ids进行mask"""
|
||||
|
||||
# MLM任务样本的标签——若token被mask,则对应标签为词表内的索引id;若没有被mask,则标签为PAD_IDX,在计算loss时会被忽略
|
||||
mlm_label = [self.PAD_IDX] * len(token_ids)
|
||||
mask_ct = 0 # 记录已被mask的tokens数量
|
||||
|
||||
for mask_pos in candidate_mask_positions:
|
||||
# 被mask的tokens数量已达到要求
|
||||
if mask_ct >= num_mask_ids:
|
||||
break
|
||||
|
||||
new_token_id = None # 用于mask(即替换)的token id
|
||||
# 15%的tokens中的80%替换为[MASK] token
|
||||
if random.random() < self.masked_token_rate:
|
||||
new_token_id = self.MASK_IDX
|
||||
else:
|
||||
# 15%的tokens中的10%保持不变
|
||||
# 20% * 0.5 = 10%
|
||||
if random.random() < self.masked_token_unchanged_rate:
|
||||
new_token_id = token_ids[mask_pos]
|
||||
# 15%的tokens中的最后10%随机替换为一个词表内的token
|
||||
else:
|
||||
new_token_id = random.randint(0, len(self.vocab.itos) - 1)
|
||||
|
||||
# 保存原索引并mask
|
||||
mlm_label[mask_pos] = token_ids[mask_pos]
|
||||
token_ids[mask_pos] = new_token_id
|
||||
mask_ct += 1
|
||||
|
||||
return token_ids, mlm_label
|
||||
|
||||
def get_masked_sample(self, token_ids):
|
||||
"""
|
||||
对token_ids进行mask处理
|
||||
:param token_ids: e.g. `[101, 1031, 4895, 2243, 1033, 10029, 2000, 2624, 1031,....]`
|
||||
:return mlm_id_seq: `[101, 1031, 103, 2243, 1033, 10029, 2000, 103, 1031, ...]`
|
||||
mlm_label: `[ 0, 0, 4895, 0, 0, 0, 0, 2624, 0,...]`
|
||||
"""
|
||||
candidate_mask_positions = [] # 候选mask位置
|
||||
for pos, id in enumerate(token_ids):
|
||||
# 在MLM任务中,不会mask特殊token
|
||||
if id in [self.CLS_IDX, self.SEP_IDX]:
|
||||
continue
|
||||
candidate_mask_positions.append(pos) # 例如[2, 3, 4, 5, ....]
|
||||
random.shuffle(candidate_mask_positions) # 打乱候选mask位置
|
||||
|
||||
# 计算需要被mask的tokens数量,BERT模型中的默认mask比例是15%
|
||||
num_mask_ids = max(1, round(len(token_ids) * self.masked_rate))
|
||||
logging.debug(f"## 被Mask的tokens数量为:{num_mask_ids}")
|
||||
|
||||
mlm_id_seq, mlm_label = self.masking_tokens(
|
||||
token_ids, candidate_mask_positions, num_mask_ids
|
||||
)
|
||||
|
||||
return mlm_id_seq, mlm_label
|
||||
|
||||
@process_cache(
|
||||
unique_keys=[
|
||||
"max_sen_len",
|
||||
"masked_rate",
|
||||
"masked_token_rate",
|
||||
"masked_token_unchanged_rate",
|
||||
"random_state",
|
||||
]
|
||||
)
|
||||
def data_process(self, filepath=None):
|
||||
"""构造NSP和MLM两个预训练任务接收格式的样本"""
|
||||
|
||||
paragraphs = self.format_data(filepath) # 格式化原始数据
|
||||
data = [] # 每个元素为一个样本,包括Masked索引序列、token_type_id序列、MLM、NSP任务的标签
|
||||
max_len = 0 # 保存最长序列长度
|
||||
|
||||
desc = f"## 正在处理NSP和MLM预训练数据集 {filepath.split(os.sep)[-1]}"
|
||||
for para in tqdm(paragraphs, ncols=80, desc=desc): # 遍历每个段落
|
||||
for i in range(len(para) - 1): # 遍历每个句子
|
||||
# 生成一条句子对样本及标签
|
||||
sen, next_sen, is_next = self.get_next_sentence_sample(
|
||||
para[i], para[i + 1], paragraphs
|
||||
)
|
||||
logging.debug(f"## 当前句子文本:{sen}")
|
||||
logging.debug(f"## 下一句文本:{next_sen}")
|
||||
logging.debug(f"## 句子对标签:{is_next}")
|
||||
# 下一句为空或者只有一个字符,舍弃
|
||||
if len(next_sen) < 2:
|
||||
logging.warning(
|
||||
f"句子 '{sen}' 的下一句 '{next_sen}' 为空,应舍弃,此时NSP标签为:{is_next}"
|
||||
)
|
||||
continue
|
||||
|
||||
# 分词、转换为索引序列
|
||||
id_seq1 = [self.vocab[token] for token in self.tokenizer(sen)]
|
||||
id_seq2 = [self.vocab[token] for token in self.tokenizer(next_sen)]
|
||||
# 拼接两个句子的索引序列,并加上[CLS]、[SEP] token
|
||||
id_seq = [self.CLS_IDX] + id_seq1 + [self.SEP_IDX] + id_seq2
|
||||
|
||||
# BERT模型最大支持512个token的序列,若超过,则截断
|
||||
if len(id_seq) > self.max_position_embeddings - 1:
|
||||
id_seq = id_seq[: self.max_position_embeddings - 1]
|
||||
id_seq += [self.SEP_IDX]
|
||||
assert len(id_seq) <= self.max_position_embeddings
|
||||
|
||||
# 创建token_type_id序列,用于表示token所在序列
|
||||
seg1 = [0] * (len(id_seq1) + 2) # 起始[CLS]和中间的[SEP]两个token属于第一个序列
|
||||
seg2 = [1] * (len(id_seq) - len(seg1)) # 末尾的[SEP]token则属于第二个序列
|
||||
seg = seg1 + seg2
|
||||
assert len(seg) == len(id_seq)
|
||||
|
||||
logging.debug(
|
||||
f"## Mask之前tokens:{[self.vocab.itos[id] for id in id_seq]}"
|
||||
)
|
||||
logging.debug(f"## Mask之前token ids:{id_seq}")
|
||||
logging.debug(f"## segment ids:{seg},序列长度:{len(seg)}")
|
||||
|
||||
# 对token索引序列进行mask操作,生成Masked序列样本及标签
|
||||
mlm_id_seq, mlm_label = self.get_masked_sample(id_seq)
|
||||
logging.debug(
|
||||
f"## Mask之后tokens:{[self.vocab.itos[id] for id in mlm_id_seq]}"
|
||||
)
|
||||
logging.debug(f"## Mask之后token ids:{mlm_id_seq}")
|
||||
logging.debug(f"## Mask之后labels:{mlm_label}")
|
||||
logging.debug("=" * 20)
|
||||
|
||||
id_seq = torch.tensor(mlm_id_seq, dtype=torch.long)
|
||||
seg = torch.tensor(seg, dtype=torch.long)
|
||||
mlm_label = torch.tensor(mlm_label, dtype=torch.long)
|
||||
nsp_label = torch.tensor(int(is_next), dtype=torch.long)
|
||||
max_len = max(max_len, id_seq.size(0))
|
||||
data.append([id_seq, seg, mlm_label, nsp_label])
|
||||
|
||||
return {"data": data, "max_len": max_len}
|
||||
|
||||
def generate_batch(self, data_batch):
|
||||
"""
|
||||
对每个批次中的样本进行处理的函数,将作为一个参数传入DataLoader的构造函数
|
||||
:param data_batch: 一个批次的数据
|
||||
"""
|
||||
b_id_seqs, b_segs, b_mlm_labels, b_nsp_labels = [], [], [], []
|
||||
|
||||
# 遍历一个批次内的样本,取出索引序列、token_type_id序列和MLM、NSP任务的样本标签
|
||||
for id_seq, seg, mlm_label, nsp_label in data_batch:
|
||||
b_id_seqs.append(id_seq)
|
||||
b_segs.append(seg)
|
||||
b_mlm_labels.append(mlm_label)
|
||||
b_nsp_labels.append(nsp_label)
|
||||
|
||||
# 填充
|
||||
# #[max_sen_len, batch_size]
|
||||
b_id_seqs = pad_sequence(
|
||||
b_id_seqs,
|
||||
padding_value=self.PAD_IDX,
|
||||
max_len=self.max_sen_len,
|
||||
batch_first=False,
|
||||
)
|
||||
|
||||
b_segs = pad_sequence(
|
||||
b_segs,
|
||||
padding_value=self.PAD_IDX,
|
||||
max_len=self.max_sen_len,
|
||||
batch_first=False,
|
||||
)
|
||||
|
||||
b_mlm_labels = pad_sequence(
|
||||
b_mlm_labels,
|
||||
padding_value=self.PAD_IDX,
|
||||
max_len=self.max_sen_len,
|
||||
batch_first=False,
|
||||
)
|
||||
|
||||
# 生成Padding mask
|
||||
# #[batch_size, max_sen_len]
|
||||
b_mask = (b_id_seqs == self.PAD_IDX).transpose(0, 1)
|
||||
|
||||
# #[batch_size, ]
|
||||
b_nsp_labels = torch.tensor(b_nsp_labels, dtype=torch.long)
|
||||
|
||||
return b_id_seqs, b_segs, b_mask, b_mlm_labels, b_nsp_labels
|
||||
|
||||
def data_loader(
|
||||
self,
|
||||
train_filepath=None,
|
||||
val_filepath=None,
|
||||
test_filepath=None,
|
||||
only_test=False,
|
||||
):
|
||||
"""
|
||||
创建DataLoader
|
||||
:param only_test: 是否只返回测试集
|
||||
"""
|
||||
|
||||
test_data = self.data_process(filepath=test_filepath)["data"]
|
||||
test_loader = DataLoader(
|
||||
test_data,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=False, # 测试集不打乱
|
||||
collate_fn=self.generate_batch,
|
||||
)
|
||||
|
||||
if only_test:
|
||||
logging.info(f"## 成功返回测试集,包含样本{len(test_loader.dataset)}个")
|
||||
return test_loader
|
||||
|
||||
tmp_data = self.data_process(filepath=train_filepath)
|
||||
train_data, max_len = tmp_data["data"], tmp_data["max_len"]
|
||||
if self.max_sen_len == "same":
|
||||
self.max_sen_len = max_len
|
||||
|
||||
train_loader = DataLoader(
|
||||
train_data,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=self.is_sample_shuffle,
|
||||
collate_fn=self.generate_batch,
|
||||
)
|
||||
|
||||
val_data = self.data_process(filepath=val_filepath)["data"]
|
||||
val_loader = DataLoader(
|
||||
val_data,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=False, # 验证集不打乱
|
||||
collate_fn=self.generate_batch,
|
||||
)
|
||||
logging.info(
|
||||
f"## 成功返回训练集样本{len(train_loader.dataset)}个,验证集样本{len(val_loader.dataset)}个,"
|
||||
f"测试集样本{len(test_loader.dataset)}个"
|
||||
)
|
||||
|
||||
return train_loader, val_loader, test_loader
|
||||
|
||||
def get_inference_samples(self, sentences=None, masked=False):
|
||||
"""
|
||||
制作推理阶段输入模型的样本
|
||||
:param sentences: 列表,每个元素表示一个文本段落
|
||||
:param masked: 传入的句子是否已被mask
|
||||
"""
|
||||
|
||||
# sentences可能是单个文本段落(字符串),需要转换为列表
|
||||
if not isinstance(sentences, list):
|
||||
sentences = [sentences]
|
||||
|
||||
mask_token = self.vocab.itos[self.MASK_IDX] # [MASK]
|
||||
b_id_seq = [] # 保存所有样本句子的索引序列
|
||||
b_mask_pos = [] # 保存所有样本句子中的mask_token位置
|
||||
|
||||
for sentence in sentences:
|
||||
# 分词,注意:推理阶段,只需把段落中所有token分开即可,不用考虑上下句关系
|
||||
token_seq = self.tokenizer(sentence)
|
||||
|
||||
# 传入的句子没有被mask,则执行mask
|
||||
if not masked:
|
||||
# 候选的mask位置
|
||||
candidate_mask_positions = [pos for pos in range(len(token_seq))]
|
||||
random.shuffle(candidate_mask_positions) # 打乱,以实现随机mask
|
||||
# 和训练时设置的mask比例15%保持一致
|
||||
num_mask_tokens = max(1, round(len(token_seq) * self.masked_rate))
|
||||
# 执行mask
|
||||
for pos in candidate_mask_positions[:num_mask_tokens]:
|
||||
token_seq[pos] = mask_token
|
||||
|
||||
# 转换为索引序列
|
||||
id_seq = [self.vocab[token] for token in token_seq]
|
||||
# 加上[CLS]和[SEP] tokens
|
||||
id_seq = [self.CLS_IDX] + id_seq + [self.SEP_IDX]
|
||||
|
||||
# 得到被mask的token位置(包含[CLS]、[SEP] token在内的序列内位置)
|
||||
b_mask_pos.append(self.get_mask_pos(id_seq))
|
||||
|
||||
b_id_seq.append(torch.tensor(id_seq, dtype=torch.long))
|
||||
|
||||
# 填充,按一个批次内的最长序列长度填充
|
||||
b_id_seq = pad_sequence(
|
||||
b_id_seq,
|
||||
padding_value=self.PAD_IDX,
|
||||
max_len=None,
|
||||
batch_first=False,
|
||||
)
|
||||
|
||||
b_mask = (b_id_seq == self.PAD_IDX).transpose(0, 1)
|
||||
|
||||
return b_id_seq, b_mask_pos, b_mask
|
||||
|
||||
def get_mask_pos(self, token_ids):
|
||||
"""返回token_ids中[MASK] token所在的位置"""
|
||||
|
||||
mask_positions = []
|
||||
for pos, id in enumerate(token_ids):
|
||||
if id == self.MASK_IDX:
|
||||
mask_positions.append(pos)
|
||||
return mask_positions
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
from ast import arg
|
||||
import os
|
||||
import time
|
||||
import logging
|
||||
|
|
@ -66,36 +67,55 @@ def pad_sequence(sequences, padding_value=0, max_len=None, batch_first=False):
|
|||
return out_tensor
|
||||
|
||||
|
||||
def cache_decorator(func):
|
||||
def process_cache(unique_keys=None):
|
||||
"""
|
||||
修饰器——缓存token转换为索引的结果
|
||||
数据预处理结果缓存修饰器
|
||||
:param unique_key: 相关数据集构造类中的成员变量,用于区分缓存结果
|
||||
"""
|
||||
|
||||
def wrapper(*args, **kwargs):
|
||||
filepath = kwargs["filepath"] # 文件路径
|
||||
filename = "".join(filepath.split(os.sep)[-1].split(".")[:-1]) # 文件名(不包含拓展名)
|
||||
filedir = f"{os.sep}".join(filepath.split(os.sep)[:-1]) # 文件目录
|
||||
if unique_keys is None:
|
||||
raise ValueError(
|
||||
"`unique_key`不能为空,需指定为相关数据集构造类的成员变量,如['max_sen_len', 'masked_rate', ...]"
|
||||
)
|
||||
|
||||
cache_filename = f"cache_{filename}_token2idx.pt"
|
||||
cache_path = os.path.join(filedir, cache_filename)
|
||||
def cache_decorator(func):
|
||||
def wrapper(*args, **kwargs):
|
||||
logging.info(f"## 预处理缓存文件的关键字为:{unique_keys}")
|
||||
filepath = kwargs["filepath"] # 文件路径
|
||||
filename = "_".join(
|
||||
filepath.split(os.sep)[-1].split(".")[:-1]
|
||||
) # 文件名(不包含拓展名)
|
||||
filedir = f"{os.sep}".join(filepath.split(os.sep)[:-1]) # 文件目录
|
||||
|
||||
start_time = time.time()
|
||||
if not os.path.exists(cache_path):
|
||||
logging.info(f"缓存文件 {cache_path} 不存在,处理数据集并缓存!")
|
||||
data = func(*args, **kwargs) # token转换为索引
|
||||
with open(cache_path, "wb") as f:
|
||||
torch.save(data, f) # 缓存
|
||||
else:
|
||||
logging.info(f"缓存文件 {cache_path} 存在,载入缓存!")
|
||||
with open(cache_path, "rb") as f:
|
||||
data = torch.load(f)
|
||||
end_time = time.time()
|
||||
obj = args[0] # 获取对象,因为data_process()的第1个参数为self,即对象本身
|
||||
cache_filename = f"cache_{filename}_" # 缓存文件名
|
||||
# 根据unique_keys和对应值,更新缓存文件名
|
||||
for key in unique_keys:
|
||||
key_abbr = "".join(
|
||||
[part[0] for part in key.split("_")]
|
||||
) # 生成key的简略写法,避免缓存文件名过长
|
||||
cache_filename += f"{key_abbr}{obj.__dict__[key]}_"
|
||||
cache_filepath = os.path.join(filedir, cache_filename[:-1] + ".pt")
|
||||
|
||||
logging.info(f"数据预处理一共耗时{(end_time - start_time):.3f}s")
|
||||
start_time = time.time()
|
||||
if not os.path.exists(cache_filepath):
|
||||
logging.info(f"缓存文件 {cache_filepath} 不存在,处理数据集并缓存!")
|
||||
data = func(*args, **kwargs) # token转换为索引
|
||||
with open(cache_filepath, "wb") as f:
|
||||
torch.save(data, f) # 缓存
|
||||
else:
|
||||
logging.info(f"缓存文件 {cache_filepath} 存在,载入缓存!")
|
||||
with open(cache_filepath, "rb") as f:
|
||||
data = torch.load(f)
|
||||
end_time = time.time()
|
||||
|
||||
return data
|
||||
logging.info(f"数据预处理一共耗时{(end_time - start_time):.3f}s")
|
||||
|
||||
return wrapper
|
||||
return data
|
||||
|
||||
return wrapper
|
||||
|
||||
return cache_decorator
|
||||
|
||||
|
||||
class LoadSenClsDataset:
|
||||
|
|
@ -139,8 +159,8 @@ class LoadSenClsDataset:
|
|||
self.SEP_IDX = self.vocab["[SEP]"]
|
||||
self.is_sample_shuffle = is_sample_shuffle
|
||||
|
||||
@cache_decorator
|
||||
def token_to_idx(self, filepath=None):
|
||||
@process_cache(unique_keys=["max_sen_len"])
|
||||
def data_process(self, filepath=None):
|
||||
"""
|
||||
将token序列转换为索引序列,并返回最长序列长度
|
||||
"""
|
||||
|
|
@ -181,7 +201,7 @@ class LoadSenClsDataset:
|
|||
创建DataLoader
|
||||
:param only_test: 是否只返回测试集
|
||||
"""
|
||||
test_data, _ = self.token_to_idx(filepath=test_filepath)
|
||||
test_data, _ = self.data_process(filepath=test_filepath)
|
||||
test_loader = DataLoader(
|
||||
test_data,
|
||||
batch_size=self.batch_size,
|
||||
|
|
@ -191,7 +211,7 @@ class LoadSenClsDataset:
|
|||
if only_test:
|
||||
return test_loader
|
||||
|
||||
train_data, max_len = self.token_to_idx(filepath=train_filepath)
|
||||
train_data, max_len = self.data_process(filepath=train_filepath)
|
||||
if self.max_sen_len == "same":
|
||||
self.max_sen_len = max_len
|
||||
|
||||
|
|
@ -202,7 +222,7 @@ class LoadSenClsDataset:
|
|||
collate_fn=self.generate_batch,
|
||||
)
|
||||
|
||||
val_data, _ = self.token_to_idx(filepath=val_filepath)
|
||||
val_data, _ = self.data_process(filepath=val_filepath)
|
||||
val_loader = DataLoader(
|
||||
val_data,
|
||||
batch_size=self.batch_size,
|
||||
|
|
@ -243,9 +263,9 @@ class LoadPairSenClsDataset(LoadSenClsDataset):
|
|||
super().__init__(**kwargs)
|
||||
pass
|
||||
|
||||
# 重载父类LoadSenClsDataset中的token_to_idx和generate_batch方法
|
||||
@cache_decorator
|
||||
def token_to_idx(self, filepath=None):
|
||||
# 覆盖父类LoadSenClsDataset中的data_process和generate_batch方法
|
||||
@process_cache(unique_keys=["max_sen_len"])
|
||||
def data_process(self, filepath=None):
|
||||
"""
|
||||
将token序列转换为索引序列,并返回最长序列长度
|
||||
"""
|
||||
|
|
|
|||
Loading…
Reference in New Issue