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 +
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]
)
]