From 6c171fcea0dd5c60abb70527a9dda88d5cb74c16 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 23 Dec 2023 12:59:57 +0900 Subject: [PATCH] Support V210 train and fix some bugs and improve --- README.md | 14 + configs/config.json | 912 +--------------------------------- data_utils.py | 4 + emo_gen.py | 155 ++++++ oldVersion/V210/data_utils.py | 414 +++++++++++++++ train_ms.py | 19 + train_ms_V210.py | 710 ++++++++++++++++++++++++++ utils.py | 37 +- 8 files changed, 1365 insertions(+), 900 deletions(-) create mode 100644 emo_gen.py create mode 100644 oldVersion/V210/data_utils.py create mode 100644 train_ms_V210.py diff --git a/README.md b/README.md index 1559afd..e9bb151 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,17 @@ +# 変更点 +- `keep_ckpts`のバグを修正 +- 圧縮モデルを保存する設定を追加: `config.json`の`save_compressed_models`を`true`にすると、modelの保存時に圧縮モデルも保存される(これと`keep_ckpts`を`1`とかにしとけばかなり容量の節約に) +- Ver 2.1での学習をサポート(`train_ms_V210.py`) + +# TODO +- [ ] Ver 2.2での学習をサポート +- [ ] Ver 2.1, 2.2での学習でのbf16対応 +- [ ] 推論のWebUIでのバージョンに応じた感情指定のサポート +- [ ] より良い推論WebUI? +- [ ] 学習のWebUI? + +以下本家のREADME.md +
LOGO diff --git a/configs/config.json b/configs/config.json index 6f3c5cc..b8b996f 100644 --- a/configs/config.json +++ b/configs/config.json @@ -2,14 +2,12 @@ "train": { "log_interval": 200, "eval_interval": 1000, + "save_compressed_models": true, "seed": 42, "epochs": 1000, "learning_rate": 0.0002, - "betas": [ - 0.8, - 0.99 - ], - "eps": 1e-09, + "betas": [0.8, 0.99], + "eps": 1e-9, "batch_size": 16, "bf16_run": false, "lr_decay": 0.99995, @@ -38,859 +36,7 @@ "mel_fmax": null, "add_blank": true, "n_speakers": 850, - "cleaned_text": true, - "spk2id": { - "派蒙_ZH": 0, - "纳西妲_ZH": 1, - "凯亚_ZH": 2, - "阿贝多_ZH": 3, - "温迪_ZH": 4, - "枫原万叶_ZH": 5, - "钟离_ZH": 6, - "荒泷一斗_ZH": 7, - "八重神子_ZH": 8, - "艾尔海森_ZH": 9, - "提纳里_ZH": 10, - "迪希雅_ZH": 11, - "卡维_ZH": 12, - "宵宫_ZH": 13, - "那维莱特_ZH": 14, - "莱依拉_ZH": 15, - "赛诺_ZH": 16, - "莫娜_ZH": 17, - "诺艾尔_ZH": 18, - "托马_ZH": 19, - "凝光_ZH": 20, - "林尼_ZH": 21, - "北斗_ZH": 22, - "柯莱_ZH": 23, - "神里绫华_ZH": 24, - "可莉_ZH": 25, - "芭芭拉_ZH": 26, - "雷电将军_ZH": 27, - "娜维娅_ZH": 28, - "芙宁娜_ZH": 29, - "珊瑚宫心海_ZH": 30, - "鹿野院平藏_ZH": 31, - "迪奥娜_ZH": 32, - "琴_ZH": 33, - "五郎_ZH": 34, - "班尼特_ZH": 35, - "达达利亚_ZH": 36, - "安柏_ZH": 37, - "莱欧斯利_ZH": 38, - "夜兰_ZH": 39, - "妮露_ZH": 40, - "辛焱_ZH": 41, - "丽莎_ZH": 42, - "珐露珊_ZH": 43, - "魈_ZH": 44, - "香菱_ZH": 45, - "迪卢克_ZH": 46, - "砂糖_ZH": 47, - "烟绯_ZH": 48, - "早柚_ZH": 49, - "云堇_ZH": 50, - "刻晴_ZH": 51, - "重云_ZH": 52, - "优菈_ZH": 53, - "胡桃_ZH": 54, - "流浪者_ZH": 55, - "久岐忍_ZH": 56, - "神里绫人_ZH": 57, - "甘雨_ZH": 58, - "戴因斯雷布_ZH": 59, - "菲谢尔_ZH": 60, - "白术_ZH": 61, - "行秋_ZH": 62, - "九条裟罗_ZH": 63, - "夏洛蒂_ZH": 64, - "雷泽_ZH": 65, - "申鹤_ZH": 66, - "荧_ZH": 67, - "空_ZH": 68, - "迪娜泽黛_ZH": 69, - "凯瑟琳_ZH": 70, - "多莉_ZH": 71, - "坎蒂丝_ZH": 72, - "琳妮特_ZH": 73, - "萍姥姥_ZH": 74, - "罗莎莉亚_ZH": 75, - "埃德_ZH": 76, - "爱贝尔_ZH": 77, - "伊迪娅_ZH": 78, - "留云借风真君_ZH": 79, - "绮良良_ZH": 80, - "陌生人_ZH": 81, - "七七_ZH": 82, - "式大将_ZH": 83, - "瑶瑶_ZH": 84, - "奥兹_ZH": 85, - "菲米尼_ZH": 86, - "米卡_ZH": 87, - "哲平_ZH": 88, - "浮游水蕈兽·元素生命_ZH": 89, - "大肉丸_ZH": 90, - "托克_ZH": 91, - "蒂玛乌斯_ZH": 92, - "昆钧_ZH": 93, - "欧菲妮_ZH": 94, - "塞琉斯_ZH": 95, - "仆人_ZH": 96, - "迈勒斯_ZH": 97, - "希格雯_ZH": 98, - "阿守_ZH": 99, - "拉赫曼_ZH": 100, - "杜拉夫_ZH": 101, - "伊利亚斯_ZH": 102, - "阿晃_ZH": 103, - "旁白_ZH": 104, - "爱德琳_ZH": 105, - "埃洛伊_ZH": 106, - "德沃沙克_ZH": 107, - "玛乔丽_ZH": 108, - "塞塔蕾_ZH": 109, - "柊千里_ZH": 110, - "海芭夏_ZH": 111, - "九条镰治_ZH": 112, - "阿娜耶_ZH": 113, - "笼钓瓶一心_ZH": 114, - "回声海螺_ZH": 115, - "劳维克_ZH": 116, - "元太_ZH": 117, - "阿扎尔_ZH": 118, - "查尔斯_ZH": 119, - "阿洛瓦_ZH": 120, - "埃勒曼_ZH": 121, - "纳比尔_ZH": 122, - "莎拉_ZH": 123, - "康纳_ZH": 124, - "博来_ZH": 125, - "玛塞勒_ZH": 126, - "阿祇_ZH": 127, - "博士_ZH": 128, - "玛格丽特_ZH": 129, - "迪尔菲_ZH": 130, - "宛烟_ZH": 131, - "羽生田千鹤_ZH": 132, - "海妮耶_ZH": 133, - "旅行者_ZH": 134, - "霍夫曼_ZH": 135, - "佐西摩斯_ZH": 136, - "鹿野奈奈_ZH": 137, - "舒伯特_ZH": 138, - "天叔_ZH": 139, - "艾莉丝_ZH": 140, - "龙二_ZH": 141, - "莺儿_ZH": 142, - "嘉良_ZH": 143, - "一心传名刀_ZH": 144, - "珊瑚_ZH": 145, - "言笑_ZH": 146, - "久利须_ZH": 147, - "嘉玛_ZH": 148, - "艾文_ZH": 149, - "克洛琳德_ZH": 150, - "丹吉尔_ZH": 151, - "女士_ZH": 152, - "白老先生_ZH": 153, - "天目十五_ZH": 154, - "老孟_ZH": 155, - "巴达维_ZH": 156, - "长生_ZH": 157, - "吴船长_ZH": 158, - "拉齐_ZH": 159, - "艾伯特_ZH": 160, - "松浦_ZH": 161, - "埃泽_ZH": 162, - "阿圆_ZH": 163, - "莫塞伊思_ZH": 164, - "阿拉夫_ZH": 165, - "杜吉耶_ZH": 166, - "石头_ZH": 167, - "百闻_ZH": 168, - "波洛_ZH": 169, - "斯坦利_ZH": 170, - "博易_ZH": 171, - "迈蒙_ZH": 172, - "掇星攫辰天君_ZH": 173, - "毗伽尔_ZH": 174, - "芙卡洛斯_ZH": 175, - "恶龙_ZH": 176, - "恕筠_ZH": 177, - "知易_ZH": 178, - "克列门特_ZH": 179, - "大慈树王_ZH": 180, - "西拉杰_ZH": 181, - "上杉_ZH": 182, - "阿尔卡米_ZH": 183, - "纯水精灵_ZH": 184, - "常九爷_ZH": 185, - "沙扎曼_ZH": 186, - "田铁嘴_ZH": 187, - "克罗索_ZH": 188, - "阿巴图伊_ZH": 189, - "阿佩普_ZH": 190, - "埃尔欣根_ZH": 191, - "萨赫哈蒂_ZH": 192, - "塔杰·拉德卡尼_ZH": 193, - "安西_ZH": 194, - "陆行岩本真蕈·元素生命_ZH": 195, - "派蒙_JP": 196, - "纳西妲_JP": 197, - "凯亚_JP": 198, - "阿贝多_JP": 199, - "温迪_JP": 200, - "枫原万叶_JP": 201, - "钟离_JP": 202, - "荒泷一斗_JP": 203, - "八重神子_JP": 204, - "艾尔海森_JP": 205, - "提纳里_JP": 206, - "迪希雅_JP": 207, - "卡维_JP": 208, - "宵宫_JP": 209, - "那维莱特_JP": 210, - "莱依拉_JP": 211, - "赛诺_JP": 212, - "莫娜_JP": 213, - "诺艾尔_JP": 214, - "托马_JP": 215, - "凝光_JP": 216, - "林尼_JP": 217, - "北斗_JP": 218, - "柯莱_JP": 219, - "神里绫华_JP": 220, - "可莉_JP": 221, - "芭芭拉_JP": 222, - "雷电将军_JP": 223, - "娜维娅_JP": 224, - "芙宁娜_JP": 225, - "珊瑚宫心海_JP": 226, - "鹿野院平藏_JP": 227, - "迪奥娜_JP": 228, - "琴_JP": 229, - "五郎_JP": 230, - "班尼特_JP": 231, - "达达利亚_JP": 232, - "安柏_JP": 233, - "莱欧斯利_JP": 234, - "夜兰_JP": 235, - "妮露_JP": 236, - "辛焱_JP": 237, - "丽莎_JP": 238, - "珐露珊_JP": 239, - "魈_JP": 240, - "香菱_JP": 241, - "迪卢克_JP": 242, - "砂糖_JP": 243, - "烟绯_JP": 244, - "早柚_JP": 245, - "云堇_JP": 246, - "刻晴_JP": 247, - "重云_JP": 248, - "优菈_JP": 249, - "胡桃_JP": 250, - "流浪者_JP": 251, - "久岐忍_JP": 252, - "神里绫人_JP": 253, - "甘雨_JP": 254, - "戴因斯雷布_JP": 255, - "菲谢尔_JP": 256, - "白术_JP": 257, - "行秋_JP": 258, - "九条裟罗_JP": 259, - "夏洛蒂_JP": 260, - "雷泽_JP": 261, - "申鹤_JP": 262, - "空_JP": 263, - "荧_JP": 264, - "迪娜泽黛_JP": 265, - "凯瑟琳_JP": 266, - "多莉_JP": 267, - "坎蒂丝_JP": 268, - "琳妮特_JP": 269, - "萍姥姥_JP": 270, - "罗莎莉亚_JP": 271, - "埃德_JP": 272, - "爱贝尔_JP": 273, - "伊迪娅_JP": 274, - "留云借风真君_JP": 275, - "绮良良_JP": 276, - "陌生人_JP": 277, - "七七_JP": 278, - "式大将_JP": 279, - "瑶瑶_JP": 280, - "奥兹_JP": 281, - "菲米尼_JP": 282, - "米卡_JP": 283, - "哲平_JP": 284, - "浮游水蕈兽·元素生命_JP": 285, - "大肉丸_JP": 286, - "托克_JP": 287, - "蒂玛乌斯_JP": 288, - "昆钧_JP": 289, - "欧菲妮_JP": 290, - "塞琉斯_JP": 291, - "仆人_JP": 292, - "迈勒斯_JP": 293, - "希格雯_JP": 294, - "阿守_JP": 295, - "拉赫曼_JP": 296, - "杜拉夫_JP": 297, - "伊利亚斯_JP": 298, - "阿晃_JP": 299, - "旁白_JP": 300, - "爱德琳_JP": 301, - "埃洛伊_JP": 302, - "德沃沙克_JP": 303, - "玛乔丽_JP": 304, - "塞塔蕾_JP": 305, - "柊千里_JP": 306, - "海芭夏_JP": 307, - "九条镰治_JP": 308, - "阿娜耶_JP": 309, - "笼钓瓶一心_JP": 310, - "回声海螺_JP": 311, - "劳维克_JP": 312, - "元太_JP": 313, - "阿扎尔_JP": 314, - "查尔斯_JP": 315, - "阿洛瓦_JP": 316, - "埃勒曼_JP": 317, - "纳比尔_JP": 318, - "莎拉_JP": 319, - "康纳_JP": 320, - "博来_JP": 321, - "玛塞勒_JP": 322, - "阿祇_JP": 323, - "博士_JP": 324, - "迪尔菲_JP": 325, - "玛格丽特_JP": 326, - "宛烟_JP": 327, - "羽生田千鹤_JP": 328, - "海妮耶_JP": 329, - "霍夫曼_JP": 330, - "旅行者_JP": 331, - "佐西摩斯_JP": 332, - "舒伯特_JP": 333, - "鹿野奈奈_JP": 334, - "天叔_JP": 335, - "龙二_JP": 336, - "艾莉丝_JP": 337, - "莺儿_JP": 338, - "嘉良_JP": 339, - "珊瑚_JP": 340, - "言笑_JP": 341, - "一心传名刀_JP": 342, - "费迪南德_JP": 343, - "久利须_JP": 344, - "嘉玛_JP": 345, - "艾文_JP": 346, - "克洛琳德_JP": 347, - "丹吉尔_JP": 348, - "天目十五_JP": 349, - "女士_JP": 350, - "老孟_JP": 351, - "白老先生_JP": 352, - "舍利夫_JP": 353, - "巴达维_JP": 354, - "拉齐_JP": 355, - "长生_JP": 356, - "吴船长_JP": 357, - "艾伯特_JP": 358, - "松浦_JP": 359, - "埃泽_JP": 360, - "阿圆_JP": 361, - "阿拉夫_JP": 362, - "莫塞伊思_JP": 363, - "石头_JP": 364, - "百闻_JP": 365, - "杜吉耶_JP": 366, - "波洛_JP": 367, - "掇星攫辰天君_JP": 368, - "迈蒙_JP": 369, - "博易_JP": 370, - "诗筠_JP": 371, - "斯坦利_JP": 372, - "毗伽尔_JP": 373, - "芙卡洛斯_JP": 374, - "恶龙_JP": 375, - "小仓澪_JP": 376, - "恕筠_JP": 377, - "知易_JP": 378, - "克列门特_JP": 379, - "大慈树王_JP": 380, - "望雅_JP": 381, - "黑田_JP": 382, - "卡莉娜_JP": 383, - "马姆杜_JP": 384, - "科林斯_JP": 385, - "上杉_JP": 386, - "西拉杰_JP": 387, - "菲尔戈黛特_JP": 388, - "一平_JP": 389, - "纯水精灵_JP": 390, - "阿尔卡米_JP": 391, - "老戴_JP": 392, - "谢赫祖拜尔_JP": 393, - "沙扎曼_JP": 394, - "田铁嘴_JP": 395, - "小野寺_JP": 396, - "百识_JP": 397, - "克罗索_JP": 398, - "莱斯格_JP": 399, - "芷巧_JP": 400, - "加藤洋平_JP": 401, - "阿巴图伊_JP": 402, - "埃尔欣根_JP": 403, - "斯嘉莉_JP": 404, - "阿佩普_JP": 405, - "巫女_JP": 406, - "卡布斯_JP": 407, - "洛伦佐_JP": 408, - "萨赫哈蒂_JP": 409, - "娜德瓦_JP": 410, - "塞德娜_JP": 411, - "塔杰·拉德卡尼_JP": 412, - "绘星_JP": 413, - "泽田_JP": 414, - "安西_JP": 415, - "拉伊德_JP": 416, - "亚卡巴_JP": 417, - "有乐斋_JP": 418, - "莱昂_JP": 419, - "尤苏波夫_JP": 420, - "夏妮_JP": 421, - "埃舍尔_JP": 422, - "萨齐因_JP": 423, - "古山_JP": 424, - "自称渊上之物_JP": 425, - "丹羽_JP": 426, - "塞萨尔的日记_JP": 427, - "派蒙_EN": 428, - "纳西妲_EN": 429, - "凯亚_EN": 430, - "阿贝多_EN": 431, - "温迪_EN": 432, - "枫原万叶_EN": 433, - "钟离_EN": 434, - "荒泷一斗_EN": 435, - "八重神子_EN": 436, - "艾尔海森_EN": 437, - "提纳里_EN": 438, - "迪希雅_EN": 439, - "卡维_EN": 440, - "宵宫_EN": 441, - "莱依拉_EN": 442, - "那维莱特_EN": 443, - "赛诺_EN": 444, - "莫娜_EN": 445, - "诺艾尔_EN": 446, - "托马_EN": 447, - "凝光_EN": 448, - "林尼_EN": 449, - "北斗_EN": 450, - "柯莱_EN": 451, - "神里绫华_EN": 452, - "可莉_EN": 453, - "芭芭拉_EN": 454, - "雷电将军_EN": 455, - "娜维娅_EN": 456, - "芙宁娜_EN": 457, - "珊瑚宫心海_EN": 458, - "鹿野院平藏_EN": 459, - "迪奥娜_EN": 460, - "五郎_EN": 461, - "琴_EN": 462, - "班尼特_EN": 463, - "达达利亚_EN": 464, - "安柏_EN": 465, - "莱欧斯利_EN": 466, - "夜兰_EN": 467, - "妮露_EN": 468, - "辛焱_EN": 469, - "珐露珊_EN": 470, - "丽莎_EN": 471, - "魈_EN": 472, - "香菱_EN": 473, - "迪卢克_EN": 474, - "砂糖_EN": 475, - "烟绯_EN": 476, - "早柚_EN": 477, - "云堇_EN": 478, - "刻晴_EN": 479, - "重云_EN": 480, - "优菈_EN": 481, - "胡桃_EN": 482, - "流浪者_EN": 483, - "久岐忍_EN": 484, - "神里绫人_EN": 485, - "甘雨_EN": 486, - "戴因斯雷布_EN": 487, - "菲谢尔_EN": 488, - "白术_EN": 489, - "行秋_EN": 490, - "九条裟罗_EN": 491, - "夏洛蒂_EN": 492, - "雷泽_EN": 493, - "申鹤_EN": 494, - "荧_EN": 495, - "空_EN": 496, - "迪娜泽黛_EN": 497, - "凯瑟琳_EN": 498, - "多莉_EN": 499, - "坎蒂丝_EN": 500, - "琳妮特_EN": 501, - "萍姥姥_EN": 502, - "罗莎莉亚_EN": 503, - "埃德_EN": 504, - "爱贝尔_EN": 505, - "伊迪娅_EN": 506, - "留云借风真君_EN": 507, - "绮良良_EN": 508, - "陌生人_EN": 509, - "七七_EN": 510, - "式大将_EN": 511, - "瑶瑶_EN": 512, - "奥兹_EN": 513, - "菲米尼_EN": 514, - "米卡_EN": 515, - "哲平_EN": 516, - "浮游水蕈兽·元素生命_EN": 517, - "大肉丸_EN": 518, - "托克_EN": 519, - "蒂玛乌斯_EN": 520, - "昆钧_EN": 521, - "欧菲妮_EN": 522, - "塞琉斯_EN": 523, - "仆人_EN": 524, - "迈勒斯_EN": 525, - "希格雯_EN": 526, - "阿守_EN": 527, - "拉赫曼_EN": 528, - "杜拉夫_EN": 529, - "伊利亚斯_EN": 530, - "阿晃_EN": 531, - "旁白_EN": 532, - "爱德琳_EN": 533, - "埃洛伊_EN": 534, - "德沃沙克_EN": 535, - "玛乔丽_EN": 536, - "塞塔蕾_EN": 537, - "柊千里_EN": 538, - "海芭夏_EN": 539, - "九条镰治_EN": 540, - "阿娜耶_EN": 541, - "笼钓瓶一心_EN": 542, - "回声海螺_EN": 543, - "劳维克_EN": 544, - "元太_EN": 545, - "阿扎尔_EN": 546, - "查尔斯_EN": 547, - "阿洛瓦_EN": 548, - "埃勒曼_EN": 549, - "纳比尔_EN": 550, - "莎拉_EN": 551, - "康纳_EN": 552, - "博来_EN": 553, - "玛塞勒_EN": 554, - "阿祇_EN": 555, - "博士_EN": 556, - "迪尔菲_EN": 557, - "宛烟_EN": 558, - "玛格丽特_EN": 559, - "羽生田千鹤_EN": 560, - "海妮耶_EN": 561, - "霍夫曼_EN": 562, - "旅行者_EN": 563, - "佐西摩斯_EN": 564, - "鹿野奈奈_EN": 565, - "舒伯特_EN": 566, - "天叔_EN": 567, - "艾莉丝_EN": 568, - "龙二_EN": 569, - "莺儿_EN": 570, - "嘉良_EN": 571, - "珊瑚_EN": 572, - "费迪南德_EN": 573, - "言笑_EN": 574, - "一心传名刀_EN": 575, - "久利须_EN": 576, - "嘉玛_EN": 577, - "艾文_EN": 578, - "克洛琳德_EN": 579, - "丹吉尔_EN": 580, - "女士_EN": 581, - "天目十五_EN": 582, - "老孟_EN": 583, - "白老先生_EN": 584, - "舍利夫_EN": 585, - "巴达维_EN": 586, - "拉齐_EN": 587, - "长生_EN": 588, - "吴船长_EN": 589, - "艾伯特_EN": 590, - "松浦_EN": 591, - "埃泽_EN": 592, - "阿圆_EN": 593, - "阿拉夫_EN": 594, - "莫塞伊思_EN": 595, - "石头_EN": 596, - "百闻_EN": 597, - "杜吉耶_EN": 598, - "波洛_EN": 599, - "斯坦利_EN": 600, - "掇星攫辰天君_EN": 601, - "迈蒙_EN": 602, - "博易_EN": 603, - "诗筠_EN": 604, - "毗伽尔_EN": 605, - "慧心_EN": 606, - "芙卡洛斯_EN": 607, - "恶龙_EN": 608, - "小仓澪_EN": 609, - "恕筠_EN": 610, - "知易_EN": 611, - "克列门特_EN": 612, - "大慈树王_EN": 613, - "维多利亚_EN": 614, - "黑田_EN": 615, - "马姆杜_EN": 616, - "科林斯_EN": 617, - "上杉_EN": 618, - "西拉杰_EN": 619, - "宁禄_EN": 620, - "纯水精灵_EN": 621, - "常九爷_EN": 622, - "阿尔卡米_EN": 623, - "沙扎曼_EN": 624, - "田铁嘴_EN": 625, - "加萨尼_EN": 626, - "克罗索_EN": 627, - "星稀_EN": 628, - "莱斯格_EN": 629, - "阿巴图伊_EN": 630, - "埃尔欣根_EN": 631, - "阿佩普_EN": 632, - "萨赫哈蒂_EN": 633, - "洛伦佐_EN": 634, - "塔杰·拉德卡尼_EN": 635, - "泽田_EN": 636, - "安西_EN": 637, - "埃舍尔_EN": 638, - "三月七_ZH": 639, - "丹恒_ZH": 640, - "希儿_ZH": 641, - "娜塔莎_ZH": 642, - "希露瓦_ZH": 643, - "瓦尔特_ZH": 644, - "佩拉_ZH": 645, - "布洛妮娅_ZH": 646, - "虎克_ZH": 647, - "素裳_ZH": 648, - "克拉拉_ZH": 649, - "符玄_ZH": 650, - "白露_ZH": 651, - "杰帕德_ZH": 652, - "景元_ZH": 653, - "藿藿_ZH": 654, - "姬子_ZH": 655, - "穹_ZH": 656, - "星_ZH": 657, - "卡芙卡_ZH": 658, - "桂乃芬_ZH": 659, - "艾丝妲_ZH": 660, - "玲可_ZH": 661, - "彦卿_ZH": 662, - "托帕_ZH": 663, - "驭空_ZH": 664, - "浮烟_ZH": 665, - "停云_ZH": 666, - "镜流_ZH": 667, - "罗刹_ZH": 668, - "卢卡_ZH": 669, - "史瓦罗_ZH": 670, - "黑塔_ZH": 671, - "桑博_ZH": 672, - "伦纳德_ZH": 673, - "明曦_ZH": 674, - "银狼_ZH": 675, - "帕姆_ZH": 676, - "青雀_ZH": 677, - "乔瓦尼_ZH": 678, - "公输师傅_ZH": 679, - "晴霓_ZH": 680, - "螺丝咕姆_ZH": 681, - "阿兰_ZH": 682, - "奥列格_ZH": 683, - "丹枢_ZH": 684, - "尾巴_ZH": 685, - "寒鸦_ZH": 686, - "雪衣_ZH": 687, - "可可利亚_ZH": 688, - "青镞_ZH": 689, - "半夏_ZH": 690, - "银枝_ZH": 691, - "大毫_ZH": 692, - "霄翰_ZH": 693, - "信使_ZH": 694, - "费斯曼_ZH": 695, - "绿芙蓉_ZH": 696, - "金人会长_ZH": 697, - "维利特_ZH": 698, - "维尔德_ZH": 699, - "斯科特_ZH": 700, - "卡波特_ZH": 701, - "刃_ZH": 702, - "岩明_ZH": 703, - "浣溪_ZH": 704, - "三月七_JP": 705, - "丹恒_JP": 706, - "希儿_JP": 707, - "娜塔莎_JP": 708, - "希露瓦_JP": 709, - "瓦尔特_JP": 710, - "佩拉_JP": 711, - "布洛妮娅_JP": 712, - "虎克_JP": 713, - "素裳_JP": 714, - "克拉拉_JP": 715, - "符玄_JP": 716, - "白露_JP": 717, - "杰帕德_JP": 718, - "景元_JP": 719, - "藿藿_JP": 720, - "姬子_JP": 721, - "卡芙卡_JP": 722, - "穹_JP": 723, - "星_JP": 724, - "桂乃芬_JP": 725, - "艾丝妲_JP": 726, - "彦卿_JP": 727, - "玲可_JP": 728, - "托帕_JP": 729, - "驭空_JP": 730, - "浮烟_JP": 731, - "停云_JP": 732, - "镜流_JP": 733, - "罗刹_JP": 734, - "卢卡_JP": 735, - "史瓦罗_JP": 736, - "黑塔_JP": 737, - "桑博_JP": 738, - "伦纳德_JP": 739, - "明曦_JP": 740, - "银狼_JP": 741, - "帕姆_JP": 742, - "青雀_JP": 743, - "乔瓦尼_JP": 744, - "公输师傅_JP": 745, - "晴霓_JP": 746, - "螺丝咕姆_JP": 747, - "阿兰_JP": 748, - "奥列格_JP": 749, - "丹枢_JP": 750, - "尾巴_JP": 751, - "寒鸦_JP": 752, - "雪衣_JP": 753, - "可可利亚_JP": 754, - "青镞_JP": 755, - "半夏_JP": 756, - "银枝_JP": 757, - "大毫_JP": 758, - "霄翰_JP": 759, - "信使_JP": 760, - "费斯曼_JP": 761, - "绿芙蓉_JP": 762, - "金人会长_JP": 763, - "维利特_JP": 764, - "维尔德_JP": 765, - "斯科特_JP": 766, - "刃_JP": 767, - "卡波特_JP": 768, - "岩明_JP": 769, - "浣溪_JP": 770, - "净砚_JP": 771, - "紫月季_JP": 772, - "歌蒂_JP": 773, - "奇怪的云骑_JP": 774, - "幻胧_JP": 775, - "斯薇塔_JP": 776, - "隐书_JP": 777, - "三月七_EN": 778, - "丹恒_EN": 779, - "希儿_EN": 780, - "娜塔莎_EN": 781, - "希露瓦_EN": 782, - "瓦尔特_EN": 783, - "佩拉_EN": 784, - "布洛妮娅_EN": 785, - "虎克_EN": 786, - "素裳_EN": 787, - "克拉拉_EN": 788, - "符玄_EN": 789, - "白露_EN": 790, - "杰帕德_EN": 791, - "景元_EN": 792, - "藿藿_EN": 793, - "姬子_EN": 794, - "卡芙卡_EN": 795, - "穹_EN": 796, - "星_EN": 797, - "桂乃芬_EN": 798, - "艾丝妲_EN": 799, - "彦卿_EN": 800, - "玲可_EN": 801, - "托帕_EN": 802, - "驭空_EN": 803, - "浮烟_EN": 804, - "停云_EN": 805, - "镜流_EN": 806, - "罗刹_EN": 807, - "卢卡_EN": 808, - "史瓦罗_EN": 809, - "黑塔_EN": 810, - "桑博_EN": 811, - "伦纳德_EN": 812, - "明曦_EN": 813, - "银狼_EN": 814, - "帕姆_EN": 815, - "青雀_EN": 816, - "乔瓦尼_EN": 817, - "公输师傅_EN": 818, - "晴霓_EN": 819, - "螺丝咕姆_EN": 820, - "阿兰_EN": 821, - "奥列格_EN": 822, - "丹枢_EN": 823, - "尾巴_EN": 824, - "寒鸦_EN": 825, - "雪衣_EN": 826, - "可可利亚_EN": 827, - "青镞_EN": 828, - "半夏_EN": 829, - "银枝_EN": 830, - "大毫_EN": 831, - "霄翰_EN": 832, - "信使_EN": 833, - "费斯曼_EN": 834, - "绿芙蓉_EN": 835, - "金人会长_EN": 836, - "维利特_EN": 837, - "维尔德_EN": 838, - "刃_EN": 839, - "卡波特_EN": 840, - "岩明_EN": 841, - "浣溪_EN": 842, - "紫月季_EN": 843, - "幻胧_EN": 844, - "女声_EN": 845, - "陆景和": 846, - "莫弈": 847, - "左然": 848, - "夏彦": 849 - } + "cleaned_text": true }, "model": { "use_spk_conditioned_encoder": true, @@ -905,52 +51,24 @@ "kernel_size": 3, "p_dropout": 0.1, "resblock": "1", - "resblock_kernel_sizes": [ - 3, - 7, - 11 - ], + "resblock_kernel_sizes": [3, 7, 11], "resblock_dilation_sizes": [ - [ - 1, - 3, - 5 - ], - [ - 1, - 3, - 5 - ], - [ - 1, - 3, - 5 - ] - ], - "upsample_rates": [ - 8, - 8, - 2, - 2, - 2 + [1, 3, 5], + [1, 3, 5], + [1, 3, 5] ], + "upsample_rates": [8, 8, 2, 2, 2], "upsample_initial_channel": 512, - "upsample_kernel_sizes": [ - 16, - 16, - 8, - 2, - 2 - ], + "upsample_kernel_sizes": [16, 16, 8, 2, 2], "n_layers_q": 3, "use_spectral_norm": false, "gin_channels": 512, "slm": { - "model": "./slm/wavlm-base-plus", - "sr": 16000, - "hidden": 768, - "nlayers": 13, - "initial_channel": 64 + "model": "./slm/wavlm-base-plus", + "sr": 16000, + "hidden": 768, + "nlayers": 13, + "initial_channel": 64 } }, "version": "2.3" diff --git a/data_utils.py b/data_utils.py index c65f0cf..04d9774 100644 --- a/data_utils.py +++ b/data_utils.py @@ -299,6 +299,10 @@ class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler): self.boundaries = boundaries self.buckets, self.num_samples_per_bucket = self._create_buckets() + logger.info(f"Bucket info: {self.num_samples_per_bucket}") + logger.info( + f"Unuseful samples: {len(self.lengths) - sum(self.num_samples_per_bucket)}" + ) self.total_size = sum(self.num_samples_per_bucket) self.num_samples = self.total_size // self.num_replicas diff --git a/emo_gen.py b/emo_gen.py new file mode 100644 index 0000000..bb73bea --- /dev/null +++ b/emo_gen.py @@ -0,0 +1,155 @@ +import argparse +import os +from pathlib import Path + +import librosa +import numpy as np +import torch +import torch.nn as nn +from torch.utils.data import Dataset +from torch.utils.data import DataLoader, Dataset +from tqdm import tqdm +from transformers import Wav2Vec2Processor +from transformers.models.wav2vec2.modeling_wav2vec2 import ( + Wav2Vec2Model, + Wav2Vec2PreTrainedModel, +) + +import utils +from config import config + + +class RegressionHead(nn.Module): + r"""Classification head.""" + + def __init__(self, config): + super().__init__() + + self.dense = nn.Linear(config.hidden_size, config.hidden_size) + self.dropout = nn.Dropout(config.final_dropout) + self.out_proj = nn.Linear(config.hidden_size, config.num_labels) + + def forward(self, features, **kwargs): + x = features + x = self.dropout(x) + x = self.dense(x) + x = torch.tanh(x) + x = self.dropout(x) + x = self.out_proj(x) + + return x + + +class EmotionModel(Wav2Vec2PreTrainedModel): + r"""Speech emotion classifier.""" + + def __init__(self, config): + super().__init__(config) + + self.config = config + self.wav2vec2 = Wav2Vec2Model(config) + self.classifier = RegressionHead(config) + self.init_weights() + + def forward( + self, + input_values, + ): + outputs = self.wav2vec2(input_values) + hidden_states = outputs[0] + hidden_states = torch.mean(hidden_states, dim=1) + logits = self.classifier(hidden_states) + + return hidden_states, logits + + +class AudioDataset(Dataset): + def __init__(self, list_of_wav_files, sr, processor): + self.list_of_wav_files = list_of_wav_files + self.processor = processor + self.sr = sr + + def __len__(self): + return len(self.list_of_wav_files) + + def __getitem__(self, idx): + wav_file = self.list_of_wav_files[idx] + audio_data, _ = librosa.load(wav_file, sr=self.sr) + processed_data = self.processor(audio_data, sampling_rate=self.sr)[ + "input_values" + ][0] + return torch.from_numpy(processed_data) + + +def process_func( + x: np.ndarray, + sampling_rate: int, + model: EmotionModel, + processor: Wav2Vec2Processor, + device: str, + embeddings: bool = False, +) -> np.ndarray: + r"""Predict emotions or extract embeddings from raw audio signal.""" + model = model.to(device) + y = processor(x, sampling_rate=sampling_rate) + y = y["input_values"][0] + y = torch.from_numpy(y).unsqueeze(0).to(device) + + # run through model + with torch.no_grad(): + y = model(y)[0 if embeddings else 1] + + # convert to numpy + y = y.detach().cpu().numpy() + + return y + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument( + "-c", "--config", type=str, default=config.bert_gen_config.config_path + ) + parser.add_argument( + "--num_processes", type=int, default=config.bert_gen_config.num_processes + ) + args, _ = parser.parse_known_args() + config_path = args.config + hps = utils.get_hparams_from_file(config_path) + + device = config.bert_gen_config.device + + model_name = "./emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim" + REPO_ID = "audeering/wav2vec2-large-robust-12-ft-emotion-msp-dim" + if not Path(model_name).joinpath("pytorch_model.bin").exists(): + utils.download_emo_models(config.mirror, REPO_ID, model_name) + + processor = Wav2Vec2Processor.from_pretrained(model_name) + model = EmotionModel.from_pretrained(model_name).to(device) + + lines = [] + with open(hps.data.training_files, encoding="utf-8") as f: + lines.extend(f.readlines()) + + with open(hps.data.validation_files, encoding="utf-8") as f: + lines.extend(f.readlines()) + + wavnames = [line.split("|")[0] for line in lines] + dataset = AudioDataset(wavnames, 16000, processor) + data_loader = DataLoader( + dataset, + batch_size=1, + shuffle=False, + num_workers=min(args.num_processes, os.cpu_count() - 1), + ) + + with torch.no_grad(): + for i, data in tqdm(enumerate(data_loader), total=len(data_loader)): + wavname = wavnames[i] + emo_path = wavname.replace(".wav", ".emo.npy") + if os.path.exists(emo_path): + continue + emb = model(data.to(device))[0].detach().cpu().numpy() + np.save(emo_path, emb) + + print("Emo vec 生成完毕!") diff --git a/oldVersion/V210/data_utils.py b/oldVersion/V210/data_utils.py new file mode 100644 index 0000000..dc07f55 --- /dev/null +++ b/oldVersion/V210/data_utils.py @@ -0,0 +1,414 @@ +import os +import random +import torch +import torch.utils.data +from tqdm import tqdm +import numpy as np +from tools.log import logger +import commons +from mel_processing import spectrogram_torch, mel_spectrogram_torch +from utils import load_wav_to_torch, load_filepaths_and_text +from text import cleaned_text_to_sequence +from config import config + +"""Multi speaker version""" + + +class TextAudioSpeakerLoader(torch.utils.data.Dataset): + """ + 1) loads audio, speaker_id, text pairs + 2) normalizes text and converts them to sequences of integers + 3) computes spectrograms from audio files. + """ + + def __init__(self, audiopaths_sid_text, hparams): + self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text) + self.max_wav_value = hparams.max_wav_value + self.sampling_rate = hparams.sampling_rate + self.filter_length = hparams.filter_length + self.hop_length = hparams.hop_length + self.win_length = hparams.win_length + self.sampling_rate = hparams.sampling_rate + self.spk_map = hparams.spk2id + self.hparams = hparams + + self.use_mel_spec_posterior = getattr( + hparams, "use_mel_posterior_encoder", False + ) + if self.use_mel_spec_posterior: + self.n_mel_channels = getattr(hparams, "n_mel_channels", 80) + + self.cleaned_text = getattr(hparams, "cleaned_text", False) + + self.add_blank = hparams.add_blank + self.min_text_len = getattr(hparams, "min_text_len", 1) + self.max_text_len = getattr(hparams, "max_text_len", 384) + + random.seed(1234) + random.shuffle(self.audiopaths_sid_text) + self._filter() + + def _filter(self): + """ + Filter text & store spec lengths + """ + # Store spectrogram lengths for Bucketing + # wav_length ~= file_size / (wav_channels * Bytes per dim) = file_size / (1 * 2) + # spec_length = wav_length // hop_length + + audiopaths_sid_text_new = [] + lengths = [] + skipped = 0 + logger.info("Init dataset...") + for _id, spk, language, text, phones, tone, word2ph in tqdm( + self.audiopaths_sid_text + ): + audiopath = f"{_id}" + if self.min_text_len <= len(phones) and len(phones) <= self.max_text_len: + phones = phones.split(" ") + tone = [int(i) for i in tone.split(" ")] + word2ph = [int(i) for i in word2ph.split(" ")] + audiopaths_sid_text_new.append( + [audiopath, spk, language, text, phones, tone, word2ph] + ) + lengths.append(os.path.getsize(audiopath) // (2 * self.hop_length)) + else: + skipped += 1 + logger.info( + "skipped: " + + str(skipped) + + ", total: " + + str(len(self.audiopaths_sid_text)) + ) + self.audiopaths_sid_text = audiopaths_sid_text_new + self.lengths = lengths + + def get_audio_text_speaker_pair(self, audiopath_sid_text): + # separate filename, speaker_id and text + audiopath, sid, language, text, phones, tone, word2ph = audiopath_sid_text + + bert, ja_bert, en_bert, phones, tone, language = self.get_text( + text, word2ph, phones, tone, language, audiopath + ) + + spec, wav = self.get_audio(audiopath) + sid = torch.LongTensor([int(self.spk_map[sid])]) + emo = torch.FloatTensor(np.load(audiopath.replace(".wav", ".emo.npy"))) + return (phones, spec, wav, sid, tone, language, bert, ja_bert, en_bert, emo) + + def get_audio(self, filename): + audio, sampling_rate = load_wav_to_torch(filename) + if sampling_rate != self.sampling_rate: + raise ValueError( + "{} {} SR doesn't match target {} SR".format( + filename, sampling_rate, self.sampling_rate + ) + ) + audio_norm = audio / self.max_wav_value + audio_norm = audio_norm.unsqueeze(0) + spec_filename = filename.replace(".wav", ".spec.pt") + if self.use_mel_spec_posterior: + spec_filename = spec_filename.replace(".spec.pt", ".mel.pt") + try: + spec = torch.load(spec_filename) + except: + if self.use_mel_spec_posterior: + spec = mel_spectrogram_torch( + audio_norm, + self.filter_length, + self.n_mel_channels, + self.sampling_rate, + self.hop_length, + self.win_length, + self.hparams.mel_fmin, + self.hparams.mel_fmax, + center=False, + ) + else: + spec = spectrogram_torch( + audio_norm, + self.filter_length, + self.sampling_rate, + self.hop_length, + self.win_length, + center=False, + ) + spec = torch.squeeze(spec, 0) + if config.train_ms_config.spec_cache: + torch.save(spec, spec_filename) + return spec, audio_norm + + def get_text(self, text, word2ph, phone, tone, language_str, wav_path): + phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str) + if self.add_blank: + phone = commons.intersperse(phone, 0) + tone = commons.intersperse(tone, 0) + language = commons.intersperse(language, 0) + for i in range(len(word2ph)): + word2ph[i] = word2ph[i] * 2 + word2ph[0] += 1 + bert_path = wav_path.replace(".wav", ".bert.pt") + try: + bert_ori = torch.load(bert_path) + assert bert_ori.shape[-1] == len(phone) + except Exception as e: + logger.warning("Bert load Failed") + logger.warning(e) + + if language_str == "ZH": + bert = bert_ori + ja_bert = torch.zeros(1024, len(phone)) + en_bert = torch.zeros(1024, len(phone)) + elif language_str == "JP": + bert = torch.zeros(1024, len(phone)) + ja_bert = bert_ori + en_bert = torch.zeros(1024, len(phone)) + elif language_str == "EN": + bert = torch.zeros(1024, len(phone)) + ja_bert = torch.zeros(1024, len(phone)) + en_bert = bert_ori + phone = torch.LongTensor(phone) + tone = torch.LongTensor(tone) + language = torch.LongTensor(language) + return bert, ja_bert, en_bert, phone, tone, language + + def get_sid(self, sid): + sid = torch.LongTensor([int(sid)]) + return sid + + def __getitem__(self, index): + return self.get_audio_text_speaker_pair(self.audiopaths_sid_text[index]) + + def __len__(self): + return len(self.audiopaths_sid_text) + + +class TextAudioSpeakerCollate: + """Zero-pads model inputs and targets""" + + def __init__(self, return_ids=False): + self.return_ids = return_ids + + def __call__(self, batch): + """Collate's training batch from normalized text, audio and speaker identities + PARAMS + ------ + batch: [text_normalized, spec_normalized, wav_normalized, sid] + """ + # Right zero-pad all one-hot text sequences to max input length + _, ids_sorted_decreasing = torch.sort( + torch.LongTensor([x[1].size(1) for x in batch]), dim=0, descending=True + ) + + max_text_len = max([len(x[0]) for x in batch]) + max_spec_len = max([x[1].size(1) for x in batch]) + max_wav_len = max([x[2].size(1) for x in batch]) + + text_lengths = torch.LongTensor(len(batch)) + spec_lengths = torch.LongTensor(len(batch)) + wav_lengths = torch.LongTensor(len(batch)) + sid = torch.LongTensor(len(batch)) + + text_padded = torch.LongTensor(len(batch), max_text_len) + tone_padded = torch.LongTensor(len(batch), max_text_len) + language_padded = torch.LongTensor(len(batch), max_text_len) + bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) + ja_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) + en_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) + emo = torch.FloatTensor(len(batch), 1024) + + spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len) + wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len) + text_padded.zero_() + tone_padded.zero_() + language_padded.zero_() + spec_padded.zero_() + wav_padded.zero_() + bert_padded.zero_() + ja_bert_padded.zero_() + en_bert_padded.zero_() + emo.zero_() + + for i in range(len(ids_sorted_decreasing)): + row = batch[ids_sorted_decreasing[i]] + + text = row[0] + text_padded[i, : text.size(0)] = text + text_lengths[i] = text.size(0) + + spec = row[1] + spec_padded[i, :, : spec.size(1)] = spec + spec_lengths[i] = spec.size(1) + + wav = row[2] + wav_padded[i, :, : wav.size(1)] = wav + wav_lengths[i] = wav.size(1) + + sid[i] = row[3] + + tone = row[4] + tone_padded[i, : tone.size(0)] = tone + + language = row[5] + language_padded[i, : language.size(0)] = language + + bert = row[6] + bert_padded[i, :, : bert.size(1)] = bert + + ja_bert = row[7] + ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert + + en_bert = row[8] + en_bert_padded[i, :, : en_bert.size(1)] = en_bert + + emo[i, :] = row[9] + + return ( + text_padded, + text_lengths, + spec_padded, + spec_lengths, + wav_padded, + wav_lengths, + sid, + tone_padded, + language_padded, + bert_padded, + ja_bert_padded, + en_bert_padded, + emo, + ) + + +class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler): + """ + Maintain similar input lengths in a batch. + Length groups are specified by boundaries. + Ex) boundaries = [b1, b2, b3] -> any batch is included either {x | b1 < length(x) <=b2} or {x | b2 < length(x) <= b3}. + + It removes samples which are not included in the boundaries. + Ex) boundaries = [b1, b2, b3] -> any x s.t. length(x) <= b1 or length(x) > b3 are discarded. + """ + + def __init__( + self, + dataset, + batch_size, + boundaries, + num_replicas=None, + rank=None, + shuffle=True, + ): + super().__init__(dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle) + self.lengths = dataset.lengths + self.batch_size = batch_size + self.boundaries = boundaries + + self.buckets, self.num_samples_per_bucket = self._create_buckets() + logger.info(f"Bucket info: {self.num_samples_per_bucket}") + logger.info( + f"Unuseful samples: {len(self.lengths) - sum(self.num_samples_per_bucket)}" + ) + self.total_size = sum(self.num_samples_per_bucket) + self.num_samples = self.total_size // self.num_replicas + + def _create_buckets(self): + buckets = [[] for _ in range(len(self.boundaries) - 1)] + for i in range(len(self.lengths)): + length = self.lengths[i] + idx_bucket = self._bisect(length) + if idx_bucket != -1: + buckets[idx_bucket].append(i) + + try: + for i in range(len(buckets) - 1, 0, -1): + if len(buckets[i]) == 0: + buckets.pop(i) + self.boundaries.pop(i + 1) + assert all(len(bucket) > 0 for bucket in buckets) + # When one bucket is not traversed + except Exception as e: + print("Bucket warning ", e) + for i in range(len(buckets) - 1, -1, -1): + if len(buckets[i]) == 0: + buckets.pop(i) + self.boundaries.pop(i + 1) + + num_samples_per_bucket = [] + for i in range(len(buckets)): + len_bucket = len(buckets[i]) + total_batch_size = self.num_replicas * self.batch_size + rem = ( + total_batch_size - (len_bucket % total_batch_size) + ) % total_batch_size + num_samples_per_bucket.append(len_bucket + rem) + return buckets, num_samples_per_bucket + + def __iter__(self): + # deterministically shuffle based on epoch + g = torch.Generator() + g.manual_seed(self.epoch) + + indices = [] + if self.shuffle: + for bucket in self.buckets: + indices.append(torch.randperm(len(bucket), generator=g).tolist()) + else: + for bucket in self.buckets: + indices.append(list(range(len(bucket)))) + + batches = [] + for i in range(len(self.buckets)): + bucket = self.buckets[i] + len_bucket = len(bucket) + if len_bucket == 0: + continue + ids_bucket = indices[i] + num_samples_bucket = self.num_samples_per_bucket[i] + + # add extra samples to make it evenly divisible + rem = num_samples_bucket - len_bucket + ids_bucket = ( + ids_bucket + + ids_bucket * (rem // len_bucket) + + ids_bucket[: (rem % len_bucket)] + ) + + # subsample + ids_bucket = ids_bucket[self.rank :: self.num_replicas] + + # batching + for j in range(len(ids_bucket) // self.batch_size): + batch = [ + bucket[idx] + for idx in ids_bucket[ + j * self.batch_size : (j + 1) * self.batch_size + ] + ] + batches.append(batch) + + if self.shuffle: + batch_ids = torch.randperm(len(batches), generator=g).tolist() + batches = [batches[i] for i in batch_ids] + self.batches = batches + + assert len(self.batches) * self.batch_size == self.num_samples + return iter(self.batches) + + def _bisect(self, x, lo=0, hi=None): + if hi is None: + hi = len(self.boundaries) - 1 + + if hi > lo: + mid = (hi + lo) // 2 + if self.boundaries[mid] < x and x <= self.boundaries[mid + 1]: + return mid + elif x <= self.boundaries[mid]: + return self._bisect(x, lo, mid) + else: + return self._bisect(x, mid + 1, hi) + else: + return -1 + + def __len__(self): + return self.num_samples // self.batch_size diff --git a/train_ms.py b/train_ms.py index c7547ae..da57c9a 100644 --- a/train_ms.py +++ b/train_ms.py @@ -390,6 +390,15 @@ def run(): scheduler_wd.step() if net_dur_disc is not None: scheduler_dur_disc.step() + if epoch == hps.train.epochs: + utils.save_compressed_models_checkpoint( + net_g, + epoch, + os.path.join( + hps.model_dir, + f"release_{global_step}.pth", + ), + ) def train_and_evaluate( @@ -731,6 +740,16 @@ def train_and_evaluate( n_ckpts_to_keep=keep_ckpts, sort_by_time=True, ) + save_compressed_models = hps.train.save_compressed_models + if save_compressed_models: + utils.save_compressed_models_checkpoint( + net_g, + epoch, + os.path.join( + hps.model_dir, + f"release_{global_step}.pth", + ), + ) global_step += 1 diff --git a/train_ms_V210.py b/train_ms_V210.py new file mode 100644 index 0000000..a2d3171 --- /dev/null +++ b/train_ms_V210.py @@ -0,0 +1,710 @@ +# flake8: noqa: E402 +import platform +import os +import torch +from torch.nn import functional as F +from torch.utils.data import DataLoader +from torch.utils.tensorboard import SummaryWriter +import torch.distributed as dist +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.cuda.amp import autocast, GradScaler +from tqdm import tqdm +import logging +from config import config +import argparse +import datetime +import gc + +logging.getLogger("numba").setLevel(logging.WARNING) +import commons +import utils +from oldVersion.V210.data_utils import ( + TextAudioSpeakerLoader, + TextAudioSpeakerCollate, + DistributedBucketSampler, +) +from oldVersion.V210.models import ( + SynthesizerTrn, + MultiPeriodDiscriminator, + DurationDiscriminator, +) +from losses import generator_loss, discriminator_loss, feature_loss, kl_loss +from mel_processing import mel_spectrogram_torch, spec_to_mel_torch +from oldVersion.V210.text.symbols import symbols + +torch.backends.cuda.matmul.allow_tf32 = True +torch.backends.cudnn.allow_tf32 = ( + True # If encontered training problem,please try to disable TF32. +) +torch.set_float32_matmul_precision("medium") +torch.backends.cuda.sdp_kernel("flash") +torch.backends.cuda.enable_flash_sdp(True) +torch.backends.cuda.enable_mem_efficient_sdp( + True +) # Not available if torch version is lower than 2.0 +torch.backends.cuda.enable_math_sdp(True) +global_step = 0 + + +def run(): + # 环境变量解析 + envs = config.train_ms_config.env + for env_name, env_value in envs.items(): + if env_name not in os.environ.keys(): + print("加载config中的配置{}".format(str(env_value))) + os.environ[env_name] = str(env_value) + print( + "加载环境变量 \nMASTER_ADDR: {},\nMASTER_PORT: {},\nWORLD_SIZE: {},\nRANK: {},\nLOCAL_RANK: {}".format( + os.environ["MASTER_ADDR"], + os.environ["MASTER_PORT"], + os.environ["WORLD_SIZE"], + os.environ["RANK"], + os.environ["LOCAL_RANK"], + ) + ) + + backend = "nccl" + if platform.system() == "Windows": + backend = "gloo" # If Windows,switch to gloo backend. + dist.init_process_group( + backend=backend, + init_method="env://", + timeout=datetime.timedelta(seconds=300), + ) # Use torchrun instead of mp.spawn + rank = dist.get_rank() + local_rank = int(os.environ["LOCAL_RANK"]) + n_gpus = dist.get_world_size() + + # 命令行/config.yml配置解析 + # hps = utils.get_hparams() + parser = argparse.ArgumentParser() + # 非必要不建议使用命令行配置,请使用config.yml文件 + parser.add_argument( + "-c", + "--config", + type=str, + default=config.train_ms_config.config_path, + help="JSON file for configuration", + ) + + parser.add_argument( + "-m", + "--model", + type=str, + help="数据集文件夹路径,请注意,数据不再默认放在/logs文件夹下。如果需要用命令行配置,请声明相对于根目录的路径", + default=config.dataset_path, + ) + args = parser.parse_args() + model_dir = os.path.join(args.model, config.train_ms_config.model) + if not os.path.exists(model_dir): + os.makedirs(model_dir) + hps = utils.get_hparams_from_file(args.config) + hps.model_dir = model_dir + # 比较路径是否相同 + if os.path.realpath(args.config) != os.path.realpath( + config.train_ms_config.config_path + ): + with open(args.config, "r", encoding="utf-8") as f: + data = f.read() + with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f: + f.write(data) + + torch.manual_seed(hps.train.seed) + torch.cuda.set_device(local_rank) + + global global_step + if rank == 0: + logger = utils.get_logger(hps.model_dir) + logger.info(hps) + utils.check_git_hash(hps.model_dir) + writer = SummaryWriter(log_dir=hps.model_dir) + writer_eval = SummaryWriter(log_dir=os.path.join(hps.model_dir, "eval")) + train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data) + train_sampler = DistributedBucketSampler( + train_dataset, + hps.train.batch_size, + [32, 300, 400, 500, 600, 700, 800, 900, 1000], + num_replicas=n_gpus, + rank=rank, + shuffle=True, + ) + collate_fn = TextAudioSpeakerCollate() + train_loader = DataLoader( + train_dataset, + num_workers=min(config.train_ms_config.num_workers, os.cpu_count() - 1), + shuffle=False, + pin_memory=True, + collate_fn=collate_fn, + batch_sampler=train_sampler, + persistent_workers=True, + prefetch_factor=4, + ) # DataLoader config could be adjusted. + if rank == 0: + eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data) + eval_loader = DataLoader( + eval_dataset, + num_workers=0, + shuffle=False, + batch_size=1, + pin_memory=True, + drop_last=False, + collate_fn=collate_fn, + ) + if ( + "use_noise_scaled_mas" in hps.model.keys() + and hps.model.use_noise_scaled_mas is True + ): + print("Using noise scaled MAS for VITS2") + mas_noise_scale_initial = 0.01 + noise_scale_delta = 2e-6 + else: + print("Using normal MAS for VITS1") + mas_noise_scale_initial = 0.0 + noise_scale_delta = 0.0 + if ( + "use_duration_discriminator" in hps.model.keys() + and hps.model.use_duration_discriminator is True + ): + print("Using duration discriminator for VITS2") + net_dur_disc = DurationDiscriminator( + hps.model.hidden_channels, + hps.model.hidden_channels, + 3, + 0.1, + gin_channels=hps.model.gin_channels if hps.data.n_speakers != 0 else 0, + ).cuda(local_rank) + if ( + "use_spk_conditioned_encoder" in hps.model.keys() + and hps.model.use_spk_conditioned_encoder is True + ): + if hps.data.n_speakers == 0: + raise ValueError( + "n_speakers must be > 0 when using spk conditioned encoder to train multi-speaker model" + ) + else: + print("Using normal encoder for VITS1") + + net_g = SynthesizerTrn( + len(symbols), + hps.data.filter_length // 2 + 1, + hps.train.segment_size // hps.data.hop_length, + n_speakers=hps.data.n_speakers, + mas_noise_scale_initial=mas_noise_scale_initial, + noise_scale_delta=noise_scale_delta, + **hps.model, + ).cuda(local_rank) + + net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank) + optim_g = torch.optim.AdamW( + filter(lambda p: p.requires_grad, net_g.parameters()), + hps.train.learning_rate, + betas=hps.train.betas, + eps=hps.train.eps, + ) + optim_d = torch.optim.AdamW( + net_d.parameters(), + hps.train.learning_rate, + betas=hps.train.betas, + eps=hps.train.eps, + ) + if net_dur_disc is not None: + optim_dur_disc = torch.optim.AdamW( + net_dur_disc.parameters(), + hps.train.learning_rate, + betas=hps.train.betas, + eps=hps.train.eps, + ) + else: + optim_dur_disc = None + net_g = DDP(net_g, device_ids=[local_rank]) + net_d = DDP(net_d, device_ids=[local_rank]) + dur_resume_lr = None + if net_dur_disc is not None: + net_dur_disc = DDP( + net_dur_disc, device_ids=[local_rank], find_unused_parameters=True + ) + + # 下载底模 + if config.train_ms_config.base["use_base_model"]: + utils.download_checkpoint( + hps.model_dir, + config.train_ms_config.base, + token=config.openi_token, + mirror=config.mirror, + ) + + try: + if net_dur_disc is not None: + _, _, dur_resume_lr, epoch_str = utils.load_checkpoint( + utils.latest_checkpoint_path(hps.model_dir, "DUR_*.pth"), + net_dur_disc, + optim_dur_disc, + skip_optimizer=hps.train.skip_optimizer + if "skip_optimizer" in hps.train + else True, + ) + _, optim_g, g_resume_lr, epoch_str = utils.load_checkpoint( + utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), + net_g, + optim_g, + skip_optimizer=hps.train.skip_optimizer + if "skip_optimizer" in hps.train + else True, + ) + _, optim_d, d_resume_lr, epoch_str = utils.load_checkpoint( + utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), + net_d, + optim_d, + skip_optimizer=hps.train.skip_optimizer + if "skip_optimizer" in hps.train + else True, + ) + if not optim_g.param_groups[0].get("initial_lr"): + optim_g.param_groups[0]["initial_lr"] = g_resume_lr + if not optim_d.param_groups[0].get("initial_lr"): + optim_d.param_groups[0]["initial_lr"] = d_resume_lr + if not optim_dur_disc.param_groups[0].get("initial_lr"): + optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr + + epoch_str = max(epoch_str, 1) + # global_step = (epoch_str - 1) * len(train_loader) + global_step = int( + utils.get_steps(utils.latest_checkpoint_path(hps.model_dir, "G_*.pth")) + ) + print( + f"******************检测到模型存在,epoch为 {epoch_str},gloabl step为 {global_step}*********************" + ) + except Exception as e: + print(e) + epoch_str = 1 + global_step = 0 + + scheduler_g = torch.optim.lr_scheduler.ExponentialLR( + optim_g, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 + ) + scheduler_d = torch.optim.lr_scheduler.ExponentialLR( + optim_d, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 + ) + if net_dur_disc is not None: + if not optim_dur_disc.param_groups[0].get("initial_lr"): + optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr + scheduler_dur_disc = torch.optim.lr_scheduler.ExponentialLR( + optim_dur_disc, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 + ) + else: + scheduler_dur_disc = None + scaler = GradScaler(enabled=hps.train.fp16_run) + + for epoch in range(epoch_str, hps.train.epochs + 1): + if rank == 0: + train_and_evaluate( + rank, + local_rank, + epoch, + hps, + [net_g, net_d, net_dur_disc], + [optim_g, optim_d, optim_dur_disc], + [scheduler_g, scheduler_d, scheduler_dur_disc], + scaler, + [train_loader, eval_loader], + logger, + [writer, writer_eval], + ) + else: + train_and_evaluate( + rank, + local_rank, + epoch, + hps, + [net_g, net_d, net_dur_disc], + [optim_g, optim_d, optim_dur_disc], + [scheduler_g, scheduler_d, scheduler_dur_disc], + scaler, + [train_loader, None], + None, + None, + ) + scheduler_g.step() + scheduler_d.step() + if net_dur_disc is not None: + scheduler_dur_disc.step() + + +def train_and_evaluate( + rank, + local_rank, + epoch, + hps, + nets, + optims, + schedulers, + scaler, + loaders, + logger, + writers, +): + net_g, net_d, net_dur_disc = nets + optim_g, optim_d, optim_dur_disc = optims + scheduler_g, scheduler_d, scheduler_dur_disc = schedulers + train_loader, eval_loader = loaders + if writers is not None: + writer, writer_eval = writers + + train_loader.batch_sampler.set_epoch(epoch) + global global_step + + net_g.train() + net_d.train() + if net_dur_disc is not None: + net_dur_disc.train() + for batch_idx, ( + x, + x_lengths, + spec, + spec_lengths, + y, + y_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + ) in enumerate(tqdm(train_loader)): + if net_g.module.use_noise_scaled_mas: + current_mas_noise_scale = ( + net_g.module.mas_noise_scale_initial + - net_g.module.noise_scale_delta * global_step + ) + net_g.module.current_mas_noise_scale = max(current_mas_noise_scale, 0.0) + x, x_lengths = x.cuda(local_rank, non_blocking=True), x_lengths.cuda( + local_rank, non_blocking=True + ) + spec, spec_lengths = spec.cuda( + local_rank, non_blocking=True + ), spec_lengths.cuda(local_rank, non_blocking=True) + y, y_lengths = y.cuda(local_rank, non_blocking=True), y_lengths.cuda( + local_rank, non_blocking=True + ) + speakers = speakers.cuda(local_rank, non_blocking=True) + tone = tone.cuda(local_rank, non_blocking=True) + language = language.cuda(local_rank, non_blocking=True) + bert = bert.cuda(local_rank, non_blocking=True) + ja_bert = ja_bert.cuda(local_rank, non_blocking=True) + en_bert = en_bert.cuda(local_rank, non_blocking=True) + emo = emo.cuda(local_rank, non_blocking=True) + + with autocast(enabled=hps.train.fp16_run): + ( + y_hat, + l_length, + attn, + ids_slice, + x_mask, + z_mask, + (z, z_p, m_p, logs_p, m_q, logs_q), + (hidden_x, logw, logw_), + loss_commit, + ) = net_g( + x, + x_lengths, + spec, + spec_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + ) + mel = spec_to_mel_torch( + spec, + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + y_mel = commons.slice_segments( + mel, ids_slice, hps.train.segment_size // hps.data.hop_length + ) + y_hat_mel = mel_spectrogram_torch( + y_hat.squeeze(1), + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.hop_length, + hps.data.win_length, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + + y = commons.slice_segments( + y, ids_slice * hps.data.hop_length, hps.train.segment_size + ) # slice + + # Discriminator + y_d_hat_r, y_d_hat_g, _, _ = net_d(y, y_hat.detach()) + with autocast(enabled=False): + loss_disc, losses_disc_r, losses_disc_g = discriminator_loss( + y_d_hat_r, y_d_hat_g + ) + loss_disc_all = loss_disc + if net_dur_disc is not None: + y_dur_hat_r, y_dur_hat_g = net_dur_disc( + hidden_x.detach(), x_mask.detach(), logw.detach(), logw_.detach() + ) + with autocast(enabled=False): + # TODO: I think need to mean using the mask, but for now, just mean all + ( + loss_dur_disc, + losses_dur_disc_r, + losses_dur_disc_g, + ) = discriminator_loss(y_dur_hat_r, y_dur_hat_g) + loss_dur_disc_all = loss_dur_disc + optim_dur_disc.zero_grad() + scaler.scale(loss_dur_disc_all).backward() + scaler.unscale_(optim_dur_disc) + commons.clip_grad_value_(net_dur_disc.parameters(), None) + scaler.step(optim_dur_disc) + + optim_d.zero_grad() + scaler.scale(loss_disc_all).backward() + scaler.unscale_(optim_d) + grad_norm_d = commons.clip_grad_value_(net_d.parameters(), None) + scaler.step(optim_d) + + with autocast(enabled=hps.train.fp16_run): + # Generator + y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat) + if net_dur_disc is not None: + y_dur_hat_r, y_dur_hat_g = net_dur_disc(hidden_x, x_mask, logw, logw_) + with autocast(enabled=False): + loss_dur = torch.sum(l_length.float()) + loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel + loss_kl = kl_loss(z_p, logs_q, m_p, logs_p, z_mask) * hps.train.c_kl + + loss_fm = feature_loss(fmap_r, fmap_g) + loss_gen, losses_gen = generator_loss(y_d_hat_g) + loss_gen_all = ( + loss_gen + loss_fm + loss_mel + loss_dur + loss_kl + loss_commit + ) + if net_dur_disc is not None: + loss_dur_gen, losses_dur_gen = generator_loss(y_dur_hat_g) + loss_gen_all += loss_dur_gen + optim_g.zero_grad() + scaler.scale(loss_gen_all).backward() + scaler.unscale_(optim_g) + grad_norm_g = commons.clip_grad_value_(net_g.parameters(), None) + scaler.step(optim_g) + scaler.update() + + if rank == 0: + if global_step % hps.train.log_interval == 0: + lr = optim_g.param_groups[0]["lr"] + losses = [loss_disc, loss_gen, loss_fm, loss_mel, loss_dur, loss_kl] + logger.info( + "Train Epoch: {} [{:.0f}%]".format( + epoch, 100.0 * batch_idx / len(train_loader) + ) + ) + logger.info([x.item() for x in losses] + [global_step, lr]) + + scalar_dict = { + "loss/g/total": loss_gen_all, + "loss/d/total": loss_disc_all, + "learning_rate": lr, + "grad_norm_d": grad_norm_d, + "grad_norm_g": grad_norm_g, + } + scalar_dict.update( + { + "loss/g/fm": loss_fm, + "loss/g/mel": loss_mel, + "loss/g/dur": loss_dur, + "loss/g/kl": loss_kl, + } + ) + scalar_dict.update( + {"loss/g/{}".format(i): v for i, v in enumerate(losses_gen)} + ) + scalar_dict.update( + {"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)} + ) + scalar_dict.update( + {"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)} + ) + + image_dict = { + "slice/mel_org": utils.plot_spectrogram_to_numpy( + y_mel[0].data.cpu().numpy() + ), + "slice/mel_gen": utils.plot_spectrogram_to_numpy( + y_hat_mel[0].data.cpu().numpy() + ), + "all/mel": utils.plot_spectrogram_to_numpy( + mel[0].data.cpu().numpy() + ), + "all/attn": utils.plot_alignment_to_numpy( + attn[0, 0].data.cpu().numpy() + ), + } + utils.summarize( + writer=writer, + global_step=global_step, + images=image_dict, + scalars=scalar_dict, + ) + + if global_step % hps.train.eval_interval == 0: + evaluate(hps, net_g, eval_loader, writer_eval) + utils.save_checkpoint( + net_g, + optim_g, + hps.train.learning_rate, + epoch, + os.path.join(hps.model_dir, "G_{}.pth".format(global_step)), + ) + utils.save_checkpoint( + net_d, + optim_d, + hps.train.learning_rate, + epoch, + os.path.join(hps.model_dir, "D_{}.pth".format(global_step)), + ) + if net_dur_disc is not None: + utils.save_checkpoint( + net_dur_disc, + optim_dur_disc, + hps.train.learning_rate, + epoch, + os.path.join(hps.model_dir, "DUR_{}.pth".format(global_step)), + ) + keep_ckpts = config.train_ms_config.keep_ckpts + if keep_ckpts > 0: + utils.clean_checkpoints( + path_to_models=hps.model_dir, + n_ckpts_to_keep=keep_ckpts, + sort_by_time=True, + ) + save_compressed_models = hps.train.save_compressed_models + if save_compressed_models: + utils.save_compressed_models_checkpoint( + net_g, + epoch, + os.path.join( + hps.model_dir, + f"release_{global_step}.pth", + ), + ) + + global_step += 1 + gc.collect() + torch.cuda.empty_cache() + if rank == 0: + logger.info("====> Epoch: {}".format(epoch)) + + +def evaluate(hps, generator, eval_loader, writer_eval): + generator.eval() + image_dict = {} + audio_dict = {} + print("Evaluating ...") + with torch.no_grad(): + for batch_idx, ( + x, + x_lengths, + spec, + spec_lengths, + y, + y_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + ) in enumerate(eval_loader): + x, x_lengths = x.cuda(), x_lengths.cuda() + spec, spec_lengths = spec.cuda(), spec_lengths.cuda() + y, y_lengths = y.cuda(), y_lengths.cuda() + speakers = speakers.cuda() + bert = bert.cuda() + ja_bert = ja_bert.cuda() + en_bert = en_bert.cuda() + tone = tone.cuda() + language = language.cuda() + emo = emo.cuda() + for use_sdp in [True, False]: + y_hat, attn, mask, *_ = generator.module.infer( + x, + x_lengths, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + emo, + y=spec, + max_len=1000, + sdp_ratio=0.0 if not use_sdp else 1.0, + ) + y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length + + mel = spec_to_mel_torch( + spec, + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + y_hat_mel = mel_spectrogram_torch( + y_hat.squeeze(1).float(), + hps.data.filter_length, + hps.data.n_mel_channels, + hps.data.sampling_rate, + hps.data.hop_length, + hps.data.win_length, + hps.data.mel_fmin, + hps.data.mel_fmax, + ) + image_dict.update( + { + f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( + y_hat_mel[0].cpu().numpy() + ) + } + ) + audio_dict.update( + { + f"gen/audio_{batch_idx}_{use_sdp}": y_hat[ + 0, :, : y_hat_lengths[0] + ] + } + ) + image_dict.update( + { + f"gt/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( + mel[0].cpu().numpy() + ) + } + ) + audio_dict.update({f"gt/audio_{batch_idx}": y[0, :, : y_lengths[0]]}) + + utils.summarize( + writer=writer_eval, + global_step=global_step, + images=image_dict, + audios=audio_dict, + audio_sampling_rate=hps.data.sampling_rate, + ) + generator.train() + + +if __name__ == "__main__": + run() diff --git a/utils.py b/utils.py index 68fd148..2e0b806 100644 --- a/utils.py +++ b/utils.py @@ -10,6 +10,7 @@ from huggingface_hub import hf_hub_download from scipy.io.wavfile import read import torch import re +from collections import OrderedDict MATPLOTLIB_FLAG = False @@ -141,6 +142,35 @@ def save_checkpoint(model, optimizer, learning_rate, iteration, checkpoint_path) ) +def save_compressed_models_checkpoint(model, iteration, checkpoint_path, ishalf=False): + logger.info( + f"Saving compressed model at iteration {iteration} to {checkpoint_path}" + ) + if hasattr(model, "module"): + state_dict = model.module.state_dict() + else: + state_dict = model.state_dict() + keys = [] + for k in state_dict: + if "enc_q" in k: + continue # noqa: E701 + keys.append(k) + + new_dict_g = ( + {k: state_dict[k].half() for k in keys} + if ishalf + else {k: state_dict[k] for k in keys} + ) + torch.save( + { + "model": new_dict_g, + "iteration": iteration, + "learning_rate": 0, + }, + checkpoint_path, + ) + + def summarize( writer, global_step, @@ -302,9 +332,10 @@ def clean_checkpoints(path_to_models="logs/44k/", n_ckpts_to_keep=2, sort_by_tim to_del = [ os.path.join(path_to_models, fn) for fn in ( - x_sorted("G")[:-n_ckpts_to_keep] - + x_sorted("D")[:-n_ckpts_to_keep] - + x_sorted("WD")[:-n_ckpts_to_keep] + x_sorted("G_")[:-n_ckpts_to_keep] + + x_sorted("D_")[:-n_ckpts_to_keep] + + x_sorted("WD_")[:-n_ckpts_to_keep] + + x_sorted("DUR_")[:-n_ckpts_to_keep] ) ]