Support V210 train and fix some bugs and improve
This commit is contained in:
14
README.md
14
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
|
||||
|
||||
<div align="center">
|
||||
|
||||
<img alt="LOGO" src="https://cdn.jsdelivr.net/gh/fishaudio/fish-diffusion@main/images/logo_512x512.png" width="256" height="256" />
|
||||
|
||||
@@ -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,43 +51,15 @@
|
||||
"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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
155
emo_gen.py
Normal file
155
emo_gen.py
Normal file
@@ -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 生成完毕!")
|
||||
414
oldVersion/V210/data_utils.py
Normal file
414
oldVersion/V210/data_utils.py
Normal file
@@ -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
|
||||
19
train_ms.py
19
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
|
||||
|
||||
|
||||
710
train_ms_V210.py
Normal file
710
train_ms_V210.py
Normal file
@@ -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()
|
||||
37
utils.py
37
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]
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user