diff --git a/examples/flux_lora_train_example01.json b/examples/flux_lora_train_example01.json index 3eef4fc..e7c8a29 100644 --- a/examples/flux_lora_train_example01.json +++ b/examples/flux_lora_train_example01.json @@ -1,55 +1,13 @@ { - "last_node_id": 134, - "last_link_id": 236, + "last_node_id": 137, + "last_link_id": 239, "nodes": [ - { - "id": 2, - "type": "FluxTrainModelSelect", - "pos": [ - 251.60342548949495, - 45.11988253765607 - ], - "size": { - "0": 430, - "1": 130 - }, - "flags": {}, - "order": 0, - "mode": 0, - "outputs": [ - { - "name": "flux_models", - "type": "TRAIN_FLUX_MODELS", - "links": [ - 179 - ], - "shape": 3 - } - ], - "properties": { - "Node name for S&R": "FluxTrainModelSelect" - }, - "widgets_values": [ - "flux1-dev-fp8.safetensors", - "flux_vae.safetensors", - "clip_l.safetensors", - "t5\\google_t5-v1_1-xxl_encoderonly-fp8_e4m3fn.safetensors" - ] - }, { "id": 38, "type": "SetNode", "pos": { "0": 1138.6033935546875, - "1": 1.119886875152588, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 1.119886875152588 }, "size": { "0": 210, @@ -58,7 +16,7 @@ "flags": { "collapsed": true }, - "order": 20, + "order": 19, "mode": 0, "inputs": [ { @@ -87,15 +45,7 @@ "type": "GetNode", "pos": { "0": 2630, - "1": 450, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 450 }, "size": { "0": 210, @@ -104,7 +54,7 @@ "flags": { "collapsed": true }, - "order": 1, + "order": 0, "mode": 0, "inputs": [], "outputs": [ @@ -126,10 +76,10 @@ { "id": 61, "type": "PreviewImage", - "pos": [ - 3707, - 610 - ], + "pos": { + "0": 3707, + "1": 610 + }, "size": { "0": 809.35400390625, "1": 458.6750793457031 @@ -144,6 +94,7 @@ "link": 90 } ], + "outputs": [], "properties": { "Node name for S&R": "PreviewImage" } @@ -153,15 +104,7 @@ "type": "GetNode", "pos": { "0": 3706.7109375, - "1": 460, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 460 }, "size": { "0": 210, @@ -170,7 +113,7 @@ "flags": { "collapsed": true }, - "order": 2, + "order": 1, "mode": 0, "inputs": [], "outputs": [ @@ -194,15 +137,7 @@ "type": "GetNode", "pos": { "0": 4765.21875, - "1": 468.45684814453125, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 468.45684814453125 }, "size": { "0": 210, @@ -211,7 +146,7 @@ "flags": { "collapsed": true }, - "order": 3, + "order": 2, "mode": 0, "inputs": [], "outputs": [ @@ -233,10 +168,10 @@ { "id": 70, "type": "VisualizeLoss", - "pos": [ - 5586, - -246 - ], + "pos": { + "0": 5586, + "1": -246 + }, "size": { "0": 254.40000915527344, "1": 198 @@ -283,10 +218,10 @@ { "id": 73, "type": "Display Any (rgthree)", - "pos": [ - 6270, - 660 - ], + "pos": { + "0": 6270, + "1": 660 + }, "size": { "0": 210, "1": 76 @@ -302,6 +237,7 @@ "dir": 3 } ], + "outputs": [], "properties": { "Node name for S&R": "Display Any (rgthree)" }, @@ -312,10 +248,10 @@ { "id": 78, "type": "AddLabel", - "pos": [ - 2023, - 1177 - ], + "pos": { + "0": 2023, + "1": 1177 + }, "size": { "0": 315, "1": 274 @@ -378,10 +314,10 @@ { "id": 79, "type": "SomethingToString", - "pos": [ - 1815, - 1177 - ], + "pos": { + "0": 1815, + "1": 1177 + }, "size": { "0": 315, "1": 82 @@ -420,10 +356,10 @@ { "id": 80, "type": "AddLabel", - "pos": [ - 2982, - 1177 - ], + "pos": { + "0": 2982, + "1": 1177 + }, "size": { "0": 315, "1": 274 @@ -486,10 +422,10 @@ { "id": 81, "type": "SomethingToString", - "pos": [ - 2774, - 1177 - ], + "pos": { + "0": 2774, + "1": 1177 + }, "size": { "0": 315, "1": 82 @@ -528,10 +464,10 @@ { "id": 82, "type": "SomethingToString", - "pos": [ - 3909, - 1177 - ], + "pos": { + "0": 3909, + "1": 1177 + }, "size": { "0": 315, "1": 82 @@ -570,10 +506,10 @@ { "id": 83, "type": "AddLabel", - "pos": [ - 4130, - 1177 - ], + "pos": { + "0": 4130, + "1": 1177 + }, "size": { "0": 315, "1": 274 @@ -636,10 +572,10 @@ { "id": 84, "type": "SomethingToString", - "pos": [ - 4963, - 1177 - ], + "pos": { + "0": 4963, + "1": 1177 + }, "size": { "0": 315, "1": 82 @@ -678,10 +614,10 @@ { "id": 85, "type": "AddLabel", - "pos": [ - 5171, - 1177 - ], + "pos": { + "0": 5171, + "1": 1177 + }, "size": { "0": 315, "1": 274 @@ -744,10 +680,10 @@ { "id": 88, "type": "Display Any (rgthree)", - "pos": [ - 1143.6034254894948, - 82.11988253765604 - ], + "pos": { + "0": 1143.6033935546875, + "1": 82.11988067626953 + }, "size": { "0": 210, "1": 76 @@ -763,6 +699,7 @@ "dir": 3 } ], + "outputs": [], "title": "Number of epochs", "properties": { "Node name for S&R": "Display Any (rgthree)" @@ -774,10 +711,10 @@ { "id": 89, "type": "UploadToHuggingFace", - "pos": [ - 5900, - 660 - ], + "pos": { + "0": 5900, + "1": 660 + }, "size": { "0": 315, "1": 178 @@ -830,10 +767,10 @@ { "id": 90, "type": "SaveImage", - "pos": [ - 5877, - -60 - ], + "pos": { + "0": 5877, + "1": -60 + }, "size": { "0": 574.23046875, "1": 414.46881103515625 @@ -848,6 +785,7 @@ "link": 138 } ], + "outputs": [], "properties": {}, "widgets_values": [ "flux_lora_loss_plot" @@ -856,10 +794,10 @@ { "id": 105, "type": "Display Any (rgthree)", - "pos": [ - 483, - -811 - ], + "pos": { + "0": 483, + "1": -811 + }, "size": { "0": 1073.7608642578125, "1": 492.8503112792969 @@ -875,6 +813,7 @@ "dir": 3 } ], + "outputs": [], "properties": { "Node name for S&R": "Display Any (rgthree)" }, @@ -882,61 +821,25 @@ "" ] }, - { - "id": 108, - "type": "TrainDatasetGeneralConfig", - "pos": [ - -1122.3832689147948, - 213.22264583435094 - ], - "size": { - "0": 315, - "1": 154 - }, - "flags": {}, - "order": 4, - "mode": 0, - "outputs": [ - { - "name": "dataset_general", - "type": "JSON", - "links": [ - 185 - ], - "slot_index": 0, - "shape": 3 - } - ], - "properties": { - "Node name for S&R": "TrainDatasetGeneralConfig" - }, - "widgets_values": [ - false, - false, - false, - 0, - false - ] - }, { "id": 109, "type": "TrainDatasetAdd", - "pos": [ - -772.383268914795, - 203.22264583435094 - ], + "pos": { + "0": -772.3832397460938, + "1": 203.22264099121094 + }, "size": { "0": 281.5897521972656, "1": 318 }, "flags": {}, - "order": 17, + "order": 20, "mode": 0, "inputs": [ { "name": "dataset_config", "type": "JSON", - "link": 185 + "link": 239 } ], "outputs": [ @@ -969,16 +872,16 @@ { "id": 111, "type": "TrainDatasetAdd", - "pos": [ - -472.38326891479505, - 203.22264583435094 - ], + "pos": { + "0": -472.3832702636719, + "1": 203.22264099121094 + }, "size": { "0": 267.5897521972656, "1": 318 }, "flags": {}, - "order": 21, + "order": 22, "mode": 0, "inputs": [ { @@ -1017,10 +920,10 @@ { "id": 112, "type": "TrainDatasetAdd", - "pos": [ - -172.3832689147949, - 203.22264583435094 - ], + "pos": { + "0": -172.38327026367188, + "1": 203.22264099121094 + }, "size": { "0": 259.5897521972656, "1": 318 @@ -1065,17 +968,19 @@ { "id": 113, "type": "Note", - "pos": [ - -732, - 63 - ], + "pos": { + "0": -732, + "1": 63 + }, "size": { "0": 462.68292236328125, "1": 79.98078918457031 }, "flags": {}, - "order": 5, + "order": 3, "mode": 0, + "inputs": [], + "outputs": [], "properties": { "text": "" }, @@ -1088,17 +993,19 @@ { "id": 115, "type": "Note", - "pos": [ - 248.60342548949495, - -89.88011746234395 - ], + "pos": { + "0": 248.60342407226562, + "1": -89.88011932373047 + }, "size": { "0": 462.68292236328125, "1": 79.98078918457031 }, "flags": {}, - "order": 6, + "order": 4, "mode": 0, + "inputs": [], + "outputs": [], "properties": { "text": "" }, @@ -1111,16 +1018,16 @@ { "id": 117, "type": "ImageConcatFromBatch", - "pos": [ - 6690, - 410 - ], + "pos": { + "0": 6690, + "1": 410 + }, "size": { "0": 315, "1": 106 }, "flags": {}, - "order": 22, + "order": 21, "mode": 0, "inputs": [ { @@ -1160,16 +1067,16 @@ { "id": 119, "type": "ImageBatchMulti", - "pos": [ - 6820, - 180 - ], + "pos": { + "0": 6820, + "1": 180 + }, "size": { "0": 210, "1": 142 }, "flags": {}, - "order": 19, + "order": 18, "mode": 0, "inputs": [ { @@ -1213,10 +1120,10 @@ { "id": 120, "type": "GetImageSizeAndCount", - "pos": [ - 6830, - 120 - ], + "pos": { + "0": 6830, + "1": 120 + }, "size": { "0": 210, "1": 86 @@ -1224,7 +1131,7 @@ "flags": { "collapsed": true }, - "order": 18, + "order": 17, "mode": 0, "inputs": [ { @@ -1272,15 +1179,7 @@ "type": "SetNode", "pos": { "0": 2170, - "1": 1177, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 1177 }, "size": { "0": 210, @@ -1320,15 +1219,7 @@ "type": "SetNode", "pos": { "0": 3128, - "1": 1177, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 1177 }, "size": { "0": 210, @@ -1368,15 +1259,7 @@ "type": "GetNode", "pos": { "0": 6640, - "1": 190, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 190 }, "size": { "0": 210, @@ -1385,7 +1268,7 @@ "flags": { "collapsed": true }, - "order": 7, + "order": 5, "mode": 0, "inputs": [], "outputs": [ @@ -1412,15 +1295,7 @@ "type": "GetNode", "pos": { "0": 6640, - "1": 230, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 230 }, "size": { "0": 210, @@ -1429,7 +1304,7 @@ "flags": { "collapsed": true }, - "order": 8, + "order": 6, "mode": 0, "inputs": [], "outputs": [ @@ -1455,15 +1330,7 @@ "type": "SetNode", "pos": { "0": 4278, - "1": 1177, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 1177 }, "size": { "0": 210, @@ -1504,15 +1371,7 @@ "type": "GetNode", "pos": { "0": 6650, - "1": 280, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 280 }, "size": { "0": 210, @@ -1521,7 +1380,7 @@ "flags": { "collapsed": true }, - "order": 9, + "order": 7, "mode": 0, "inputs": [], "outputs": [ @@ -1547,15 +1406,7 @@ "type": "SetNode", "pos": { "0": 5319, - "1": 1177, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 1177 }, "size": { "0": 210, @@ -1595,15 +1446,7 @@ "type": "GetNode", "pos": { "0": 6640, - "1": 330, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 330 }, "size": { "0": 210, @@ -1612,7 +1455,7 @@ "flags": { "collapsed": true }, - "order": 10, + "order": 8, "mode": 0, "inputs": [], "outputs": [ @@ -1636,17 +1479,19 @@ { "id": 131, "type": "Note", - "pos": [ - 478, - -884 - ], + "pos": { + "0": 478, + "1": -884 + }, "size": { "0": 210, "1": 58 }, "flags": {}, - "order": 11, + "order": 9, "mode": 0, + "inputs": [], + "outputs": [], "properties": { "text": "" }, @@ -1659,10 +1504,10 @@ { "id": 65, "type": "FluxTrainValidate", - "pos": [ - 4775.216642000014, - 518.4568783310547 - ], + "pos": { + "0": 4775.216796875, + "1": 518.4568481445312 + }, "size": { "0": 468.5999755859375, "1": 46 @@ -1711,10 +1556,10 @@ { "id": 46, "type": "PreviewImage", - "pos": [ - 2654, - 609 - ], + "pos": { + "0": 2654, + "1": 609 + }, "size": { "0": 850.0181274414062, "1": 452.6767578125 @@ -1729,6 +1574,7 @@ "link": 70 } ], + "outputs": [], "properties": { "Node name for S&R": "PreviewImage" } @@ -1736,10 +1582,10 @@ { "id": 66, "type": "PreviewImage", - "pos": [ - 4785, - 628 - ], + "pos": { + "0": 4785, + "1": 628 + }, "size": { "0": 850.0181274414062, "1": 452.6767578125 @@ -1754,6 +1600,7 @@ "link": 95 } ], + "outputs": [], "properties": { "Node name for S&R": "PreviewImage" } @@ -1761,17 +1608,18 @@ { "id": 37, "type": "FluxTrainValidationSettings", - "pos": [ - 775, - 18 - ], + "pos": { + "0": 775, + "1": 18 + }, "size": { "0": 315, "1": 250 }, "flags": {}, - "order": 12, + "order": 10, "mode": 0, + "inputs": [], "outputs": [ { "name": "validation_settings", @@ -1801,17 +1649,19 @@ { "id": 116, "type": "Note", - "pos": [ - 776, - -111 - ], + "pos": { + "0": 776, + "1": -111 + }, "size": { "0": 308.08209228515625, "1": 78.06562805175781 }, "flags": {}, - "order": 13, + "order": 11, "mode": 0, + "inputs": [], + "outputs": [], "properties": { "text": "" }, @@ -1824,10 +1674,10 @@ { "id": 9, "type": "PreviewImage", - "pos": [ - 1547, - 596 - ], + "pos": { + "0": 1547, + "1": 596 + }, "size": { "0": 891.4732666015625, "1": 476.6578063964844 @@ -1842,6 +1692,7 @@ "link": 8 } ], + "outputs": [], "properties": { "Node name for S&R": "PreviewImage" } @@ -1849,10 +1700,10 @@ { "id": 14, "type": "FluxTrainSave", - "pos": [ - 1988, - 256 - ], + "pos": { + "0": 1988, + "1": 256 + }, "size": { "0": 393, "1": 122 @@ -1903,10 +1754,10 @@ { "id": 8, "type": "FluxTrainValidate", - "pos": [ - 1552, - 500 - ], + "pos": { + "0": 1552, + "1": 500 + }, "size": { "0": 468.5999755859375, "1": 46 @@ -1956,15 +1807,7 @@ "type": "GetNode", "pos": { "0": 1546, - "1": 433, - "2": 0, - "3": 0, - "4": 0, - "5": 0, - "6": 0, - "7": 0, - "8": 0, - "9": 0 + "1": 433 }, "size": { "0": 277.0899353027344, @@ -1973,7 +1816,7 @@ "flags": { "collapsed": true }, - "order": 14, + "order": 12, "mode": 0, "inputs": [], "outputs": [ @@ -1995,10 +1838,10 @@ { "id": 45, "type": "FluxTrainValidate", - "pos": [ - 2640, - 500 - ], + "pos": { + "0": 2640, + "1": 500 + }, "size": { "0": 468.5999755859375, "1": 46 @@ -2046,10 +1889,10 @@ { "id": 60, "type": "FluxTrainValidate", - "pos": [ - 3716.7086387499994, - 510 - ], + "pos": { + "0": 3716.708740234375, + "1": 510 + }, "size": { "0": 468.5999755859375, "1": 46 @@ -2097,10 +1940,10 @@ { "id": 47, "type": "FluxTrainSave", - "pos": [ - 3114, - 323 - ], + "pos": { + "0": 3114, + "1": 323 + }, "size": { "0": 393, "1": 122 @@ -2150,10 +1993,10 @@ { "id": 129, "type": "AddLabel", - "pos": [ - 6937, - 60 - ], + "pos": { + "0": 6937, + "1": 60 + }, "size": { "0": 315, "1": 274 @@ -2216,10 +2059,10 @@ { "id": 62, "type": "FluxTrainSave", - "pos": [ - 4202, - 331 - ], + "pos": { + "0": 4202, + "1": 331 + }, "size": { "0": 393, "1": 122 @@ -2269,10 +2112,10 @@ { "id": 134, "type": "FluxTrainSave", - "pos": [ - 5275, - 328 - ], + "pos": { + "0": 5275, + "1": 328 + }, "size": { "0": 393, "1": 122 @@ -2322,10 +2165,10 @@ { "id": 97, "type": "VisualizeLoss", - "pos": [ - 1700, - -650 - ], + "pos": { + "0": 1700, + "1": -650 + }, "size": { "0": 303.6300048828125, "1": 198 @@ -2372,10 +2215,10 @@ { "id": 99, "type": "VisualizeLoss", - "pos": [ - 2950, - -650 - ], + "pos": { + "0": 2950, + "1": -650 + }, "size": { "0": 254.40000915527344, "1": 198 @@ -2422,10 +2265,10 @@ { "id": 101, "type": "VisualizeLoss", - "pos": [ - 4090, - -650 - ], + "pos": { + "0": 4090, + "1": -650 + }, "size": { "0": 254.40000915527344, "1": 198 @@ -2472,10 +2315,10 @@ { "id": 98, "type": "SaveImage", - "pos": [ - 1680, - -340 - ], + "pos": { + "0": 1680, + "1": -340 + }, "size": { "0": 645.9608764648438, "1": 439.37261962890625 @@ -2490,6 +2333,7 @@ "link": 161 } ], + "outputs": [], "properties": {}, "widgets_values": [ "flux_lora_loss_plot" @@ -2498,10 +2342,10 @@ { "id": 100, "type": "SaveImage", - "pos": [ - 2990, - -340 - ], + "pos": { + "0": 2990, + "1": -340 + }, "size": { "0": 574.23046875, "1": 414.46881103515625 @@ -2516,6 +2360,7 @@ "link": 163 } ], + "outputs": [], "properties": {}, "widgets_values": [ "flux_lora_loss_plot" @@ -2524,10 +2369,10 @@ { "id": 102, "type": "SaveImage", - "pos": [ - 4080, - -340 - ], + "pos": { + "0": 4080, + "1": -340 + }, "size": { "0": 574.23046875, "1": 414.46881103515625 @@ -2542,6 +2387,7 @@ "link": 165 } ], + "outputs": [], "properties": {}, "widgets_values": [ "flux_lora_loss_plot" @@ -2550,17 +2396,18 @@ { "id": 95, "type": "OptimizerConfig", - "pos": [ - 322, - 385 - ], + "pos": { + "0": 322, + "1": 385 + }, "size": { "0": 315, - "1": 243.99998474121094 + "1": 244 }, "flags": {}, - "order": 15, + "order": 13, "mode": 0, + "inputs": [], "outputs": [ { "name": "optimizer_settings", @@ -2585,52 +2432,13 @@ "" ] }, - { - "id": 114, - "type": "OptimizerConfigAdafactor", - "pos": [ - 321, - 692 - ], - "size": { - "0": 315, - "1": 316 - }, - "flags": {}, - "order": 16, - "mode": 0, - "outputs": [ - { - "name": "optimizer_settings", - "type": "ARGS", - "links": null, - "shape": 3 - } - ], - "properties": { - "Node name for S&R": "OptimizerConfigAdafactor" - }, - "widgets_values": [ - 1, - "constant", - 0, - 1, - 1, - false, - false, - false, - 1, - 5, - "" - ] - }, { "id": 74, "type": "Display Any (rgthree)", - "pos": [ - 6275, - 492 - ], + "pos": { + "0": 6275, + "1": 492 + }, "size": { "0": 358.62896728515625, "1": 76 @@ -2646,6 +2454,7 @@ "dir": 3 } ], + "outputs": [], "properties": { "Node name for S&R": "Display Any (rgthree)" }, @@ -2656,10 +2465,10 @@ { "id": 133, "type": "FluxTrainEnd", - "pos": [ - 5870, - 492 - ], + "pos": { + "0": 5870, + "1": 492 + }, "size": { "0": 317.4000244140625, "1": 98 @@ -2681,8 +2490,8 @@ "links": [ 231 ], - "shape": 3, - "slot_index": 0 + "slot_index": 0, + "shape": 3 }, { "name": "metadata", @@ -2697,8 +2506,8 @@ 230, 236 ], - "shape": 3, - "slot_index": 2 + "slot_index": 2, + "shape": 3 } ], "properties": { @@ -2713,10 +2522,10 @@ { "id": 130, "type": "SaveImage", - "pos": [ - 7132, - 121 - ], + "pos": { + "0": 7132, + "1": 121 + }, "size": { "0": 619.8221435546875, "1": 714.4110107421875 @@ -2731,6 +2540,7 @@ "link": 214 } ], + "outputs": [], "properties": {}, "widgets_values": [ "flux_lora_trainer_sheet" @@ -2739,10 +2549,10 @@ { "id": 64, "type": "FluxTrainLoop", - "pos": [ - 4770, - 330 - ], + "pos": { + "0": 4770, + "1": 330 + }, "size": { "0": 393, "1": 78 @@ -2789,10 +2599,10 @@ { "id": 59, "type": "FluxTrainLoop", - "pos": [ - 3700, - 330 - ], + "pos": { + "0": 3700, + "1": 330 + }, "size": { "0": 393, "1": 78 @@ -2824,8 +2634,8 @@ "links": [ 234 ], - "shape": 3, - "slot_index": 1 + "slot_index": 1, + "shape": 3 } ], "properties": { @@ -2840,10 +2650,10 @@ { "id": 44, "type": "FluxTrainLoop", - "pos": [ - 2630, - 330 - ], + "pos": { + "0": 2630, + "1": 330 + }, "size": { "0": 393, "1": 78 @@ -2875,8 +2685,8 @@ "links": [ 235 ], - "shape": 3, - "slot_index": 1 + "slot_index": 1, + "shape": 3 } ], "properties": { @@ -2891,10 +2701,10 @@ { "id": 4, "type": "FluxTrainLoop", - "pos": [ - 1519, - 256 - ], + "pos": { + "0": 1519, + "1": 256 + }, "size": { "0": 393, "1": 78 @@ -2926,8 +2736,8 @@ "links": [ 220 ], - "shape": 3, - "slot_index": 1 + "slot_index": 1, + "shape": 3 } ], "properties": { @@ -2939,17 +2749,143 @@ "color": "#232", "bgcolor": "#353" }, + { + "id": 135, + "type": "StringConstantMultiline", + "pos": { + "0": 319, + "1": 729 + }, + "size": { + "0": 400, + "1": 200 + }, + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "links": [ + 237 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "StringConstantMultiline" + }, + "widgets_values": [ + "cute anime girl blonde messy long hair blue eyes wearing a maid outfit with a long black dress with a gold leaf pattern and a white apron in an old dark victorian mansion with a bright window and very expensive stuff everywhere akihikoyoshida|illustration of a kitten akihikoyoshida|photograph of a turtle akihikoyoshida|portrait of a female red wizard akihikoyoshida", + true + ] + }, + { + "id": 136, + "type": "FluxTrainModelSelect", + "pos": { + "0": 251.60342407226562, + "1": 45.11988067626953 + }, + "size": { + "0": 427.607421875, + "1": 137.3937225341797 + }, + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "lora_path", + "type": "STRING", + "link": null, + "widget": { + "name": "lora_path" + } + } + ], + "outputs": [ + { + "name": "flux_models", + "type": "TRAIN_FLUX_MODELS", + "links": [ + 238 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainModelSelect" + }, + "widgets_values": [ + "flux1-dev-fp8.safetensors", + "flux_vae.safetensors", + "clip_l.safetensors", + "t5\\google_t5-v1_1-xxl_encoderonly-fp8_e4m3fn.safetensors", + "" + ] + }, + { + "id": 137, + "type": "TrainDatasetGeneralConfig", + "pos": { + "0": -1122, + "1": 203 + }, + "size": [ + 316.326628330723, + 184.85814244742187 + ], + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "reg_data_dir", + "type": "STRING", + "link": null, + "widget": { + "name": "reg_data_dir" + } + } + ], + "outputs": [ + { + "name": "dataset_general", + "type": "JSON", + "links": [ + 239 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "TrainDatasetGeneralConfig" + }, + "widgets_values": [ + false, + false, + false, + 0, + false, + false, + "" + ] + }, { "id": 107, "type": "InitFluxLoRATraining", - "pos": [ - 783, - 326 - ], - "size": [ - 449.9458725932643, - 853.0341104055544 - ], + "pos": { + "0": 783, + "1": 326 + }, + "size": { + "0": 477.3700866699219, + "1": 877.820068359375 + }, "flags": {}, "order": 24, "mode": 0, @@ -2957,7 +2893,7 @@ { "name": "flux_models", "type": "TRAIN_FLUX_MODELS", - "link": 179 + "link": 238 }, { "name": "dataset", @@ -2973,6 +2909,19 @@ "name": "resume_args", "type": "ARGS", "link": null + }, + { + "name": "block_args", + "type": "ARGS", + "link": null + }, + { + "name": "sample_prompts", + "type": "STRING", + "link": 237, + "widget": { + "name": "sample_prompts" + } } ], "outputs": [ @@ -3029,8 +2978,11 @@ "bf16", "bf16", "sdpa", - "cute anime girl blonde messy long hair blue eyes wearing a maid outfit with a long black dress with a gold leaf pattern and a white apron in an old dark victorian mansion with a bright window and very expensive stuff everywhere akihikoyoshida|illustration of a kitten akihikoyoshida|photograph of a turtle akihikoyoshida|portrait of a female red wizard akihikoyoshida", - "" + "", + "", + "disabled", + 0, + "enabled" ] } ], @@ -3235,14 +3187,6 @@ 0, "NETWORKTRAINER" ], - [ - 179, - 2, - 0, - 107, - 0, - "TRAIN_FLUX_MODELS" - ], [ 180, 95, @@ -3275,14 +3219,6 @@ 0, "*" ], - [ - 185, - 108, - 0, - 109, - 0, - "JSON" - ], [ 187, 109, @@ -3570,56 +3506,44 @@ 74, 0, "*" + ], + [ + 237, + 135, + 0, + 107, + 5, + "STRING" + ], + [ + 238, + 136, + 0, + 107, + 0, + "TRAIN_FLUX_MODELS" + ], + [ + 239, + 137, + 0, + 109, + 0, + "JSON" ] ], "groups": [ { - "title": "Train_01", + "title": "Dataset", "bounding": [ - 1439, - 120, - 1107, - 975 + -1190, + -151, + 1362, + 851 ], "color": "#3f789e", "font_size": 24, - "locked": false - }, - { - "title": "Settings and init", - "bounding": [ - 195, - -187, - 1199, - 1405 - ], - "color": "#b06634", - "font_size": 24, - "locked": false - }, - { - "title": "Train_02", - "bounding": [ - 2602, - 124, - 1046, - 975 - ], - "color": "#3f789e", - "font_size": 24, - "locked": false - }, - { - "title": "Train_03", - "bounding": [ - 3681, - 128, - 1047, - 986 - ], - "color": "#3f789e", - "font_size": 24, - "locked": false + "flags": {} }, { "title": "Train_04", @@ -3631,28 +3555,64 @@ ], "color": "#3f789e", "font_size": 24, - "locked": false + "flags": {} }, { - "title": "Dataset", + "title": "Train_03", "bounding": [ - -1190, - -151, - 1362, - 851 + 3681, + 128, + 1047, + 986 ], "color": "#3f789e", "font_size": 24, - "locked": false + "flags": {} + }, + { + "title": "Train_02", + "bounding": [ + 2602, + 124, + 1046, + 975 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "title": "Settings and init", + "bounding": [ + 195, + -187, + 1223, + 1511 + ], + "color": "#b06634", + "font_size": 24, + "flags": {} + }, + { + "title": "Train_01", + "bounding": [ + 1439, + 120, + 1107, + 975 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} } ], "config": {}, "extra": { "ds": { - "scale": 0.6830134553650705, + "scale": 0.751314800901578, "offset": [ - 177.954686248011, - 353.1723637169663 + 762.1940777239663, + 226.3834867029684 ] } }, diff --git a/examples/flux_train_example_01 b/examples/flux_train_example_01 new file mode 100644 index 0000000..a17dfa7 Binary files /dev/null and b/examples/flux_train_example_01 differ diff --git a/examples/flux_train_example_01.png b/examples/flux_train_example_01.png deleted file mode 100644 index 5be3fe2..0000000 Binary files a/examples/flux_train_example_01.png and /dev/null differ diff --git a/flux_train_comfy.py b/flux_train_comfy.py index 22527fc..b6fdc61 100644 --- a/flux_train_comfy.py +++ b/flux_train_comfy.py @@ -680,7 +680,7 @@ class FluxTrainer: else: with torch.no_grad(): # encode images to latents. images are [-1, 1] - latents = ae.encode(batch["images"]) + latents = ae.encode(batch["images"].to(ae.dtype)).to(accelerator.device, dtype=weight_dtype) # NaNが含まれていれば警告を表示し0に置き換える if torch.any(torch.isnan(latents)): diff --git a/flux_train_network_comfy.py b/flux_train_network_comfy.py index c276fae..c4b4513 100644 --- a/flux_train_network_comfy.py +++ b/flux_train_network_comfy.py @@ -4,7 +4,7 @@ import math from typing import Any import argparse from .library import flux_models, flux_train_utils, flux_utils, sd3_train_utils, strategy_base, strategy_flux, train_util -from .train_network import NetworkTrainer, clean_memory_on_device +from .train_network import NetworkTrainer, clean_memory_on_device, setup_parser from accelerate import Accelerator @@ -20,19 +20,25 @@ class FluxNetworkTrainer(NetworkTrainer): def assert_extra_args(self, args, train_dataset_group): super().assert_extra_args(args, train_dataset_group) + # sdxl_train_util.verify_sdxl_training_args(args) + + if args.fp8_base_unet: + args.fp8_base = True # if fp8_base_unet is enabled, fp8_base is also enabled for FLUX.1 + + if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs: + logger.warning( + "cache_text_encoder_outputs_to_disk is enabled, so cache_text_encoder_outputs is also enabled / cache_text_encoder_outputs_to_diskが有効になっているため、cache_text_encoder_outputsも有効になります" + ) + args.cache_text_encoder_outputs = True if args.cache_text_encoder_outputs: assert ( train_dataset_group.is_text_encoder_output_cacheable() - ), "when caching Text Encoder output, either caption_dropout_rate, shuffle_caption, token_warmup_step or caption_tag_dropout_rate cannot be used" + ), "when caching Text Encoder output, either caption_dropout_rate, shuffle_caption, token_warmup_step or caption_tag_dropout_rate cannot be used / Text Encoderの出力をキャッシュするときはcaption_dropout_rate, shuffle_caption, token_warmup_step, caption_tag_dropout_rateは使えません" - #assert ( - # args.network_train_unet_only or not args.cache_text_encoder_outputs - #), "network for Text Encoder cannot be trained with caching Text Encoder outputs" - if not args.network_train_unet_only: - logger.info( - "network for CLIP-L only will be trained. T5XXL will not be trained / CLIP-Lのネットワークのみが学習されます。T5XXLは学習されません" - ) + # prepare CLIP-L/T5XXL training flags + self.train_clip_l = not args.network_train_unet_only + self.train_t5xxl = False # default is False even if args.network_train_unet_only is False if args.max_token_length is not None: logger.warning("max_token_length is not used in Flux training") @@ -41,32 +47,60 @@ class FluxNetworkTrainer(NetworkTrainer): "split_mode and cpu_offload_checkpointing cannot be used together" ) + assert not args.split_mode or not args.cpu_offload_checkpointing, ( + "split_mode and cpu_offload_checkpointing cannot be used together" + ) + train_dataset_group.verify_bucket_reso_steps(32) # TODO check this def get_flux_model_name(self, args): return "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev" - + def load_target_model(self, args, weight_dtype, accelerator): # currently offload to cpu for some models name = self.get_flux_model_name(args) - # if we load to cpu, flux.to(fp8) takes a long time - model = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu") + + # if the file is fp8 and we are using fp8_base, we can load it as is (fp8) + loading_dtype = None if args.fp8_base else weight_dtype + + # if we load to cpu, flux.to(fp8) takes a long time, so we should load to gpu in future + model = flux_utils.load_flow_model( + name, args.pretrained_model_name_or_path, loading_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors + ) + if args.fp8_base: + # check dtype of model + if model.dtype == torch.float8_e4m3fnuz or model.dtype == torch.float8_e5m2fnuz: + raise ValueError(f"Unsupported fp8 model dtype: {model.dtype}") + elif model.dtype == torch.float8_e4m3fn or model.dtype == torch.float8_e5m2: + logger.info(f"Loaded {model.dtype} FLUX model") if args.split_mode: - model = self.prepare_split_model(model, weight_dtype, accelerator, args) + model = self.prepare_split_model(model, args, weight_dtype, accelerator) - clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu") + clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors) clip_l.eval() - # loading t5xxl to cpu takes a long time, so we should load to gpu in future - t5xxl = flux_utils.load_t5xxl(args.t5xxl, weight_dtype, "cpu") - t5xxl.eval() + # if the file is fp8 and we are using fp8_base (not unet), we can load it as is (fp8) + if args.fp8_base and not args.fp8_base_unet: + loading_dtype = None # as is + else: + loading_dtype = weight_dtype - ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu") + # loading t5xxl to cpu takes a long time, so we should load to gpu in future + t5xxl = flux_utils.load_t5xxl(args.t5xxl, loading_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors) + t5xxl.eval() + if args.fp8_base and not args.fp8_base_unet: + # check dtype of model + if t5xxl.dtype == torch.float8_e4m3fnuz or t5xxl.dtype == torch.float8_e5m2 or t5xxl.dtype == torch.float8_e5m2fnuz: + raise ValueError(f"Unsupported fp8 model dtype: {t5xxl.dtype}") + elif t5xxl.dtype == torch.float8_e4m3fn: + logger.info("Loaded fp8 T5XXL model") + + ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu", disable_mmap=args.disable_mmap_load_safetensors) return flux_utils.MODEL_VERSION_FLUX_V1, [clip_l, t5xxl], ae, model - def prepare_split_model(self, model, weight_dtype, accelerator, args): + def prepare_split_model(self, model, args, weight_dtype, accelerator): from accelerate import init_empty_weights logger.info("prepare split model") @@ -85,7 +119,13 @@ class FluxNetworkTrainer(NetworkTrainer): flux_upper.load_state_dict(sd, strict=False, assign=True) logger.info("prepare upper model") - target_dtype = torch.float8_e4m3fn if args.fp8_base else weight_dtype + if args.fp8_base: + if args.fp8_dtype and args.fp8_dtype.lower() == "e5m2": + target_dtype = torch.float8_e5m2 + else: + target_dtype = torch.float8_e4m3fn + else: + target_dtype =weight_dtype flux_upper.to(accelerator.device, dtype=target_dtype) flux_upper.eval() @@ -127,25 +167,35 @@ class FluxNetworkTrainer(NetworkTrainer): def get_text_encoding_strategy(self, args): return strategy_flux.FluxTextEncodingStrategy(apply_t5_attn_mask=args.apply_t5_attn_mask) + def post_process_network(self, args, accelerator, network, text_encoders, unet): + # check t5xxl is trained or not + self.train_t5xxl = network.train_t5xxl + + if self.train_t5xxl and args.cache_text_encoder_outputs: + raise ValueError( + "T5XXL is trained, so cache_text_encoder_outputs cannot be used / T5XXL学習時はcache_text_encoder_outputsは使用できません" + ) + def get_models_for_text_encoding(self, args, accelerator, text_encoders): if args.cache_text_encoder_outputs: - if self.is_train_text_encoder(args): + if self.train_clip_l and not self.train_t5xxl: return text_encoders[0:1] # only CLIP-L is needed for encoding because T5XXL is cached else: - return text_encoders # ignored + return None # no text encoders are needed for encoding because both are cached else: return text_encoders # both CLIP-L and T5XXL are needed for encoding def get_text_encoders_train_flags(self, args, text_encoders): - return [True, False] if self.is_train_text_encoder(args) else [False, False] + return [self.train_clip_l, self.train_t5xxl] def get_text_encoder_outputs_caching_strategy(self, args): if args.cache_text_encoder_outputs: + # if the text encoders is trained, we need tokenization, so is_partial is True return strategy_flux.FluxTextEncoderOutputsCachingStrategy( args.cache_text_encoder_outputs_to_disk, None, False, - is_partial=self.is_train_text_encoder(args), + is_partial=self.train_clip_l or self.train_t5xxl, apply_t5_attn_mask=args.apply_t5_attn_mask, ) else: @@ -166,13 +216,20 @@ class FluxNetworkTrainer(NetworkTrainer): # When TE is not be trained, it will not be prepared so we need to use explicit autocast logger.info("move text encoders to gpu") - text_encoders[0].to(accelerator.device, dtype=weight_dtype) - text_encoders[1].to(accelerator.device, dtype=weight_dtype) + text_encoders[0].to(accelerator.device, dtype=weight_dtype) # always not fp8 + text_encoders[1].to(accelerator.device) + + if text_encoders[1].dtype == torch.float8_e4m3fn: + # if we load fp8 weights, the model is already fp8, so we use it as is + self.prepare_text_encoder_fp8(1, text_encoders[1], text_encoders[1].dtype, weight_dtype) + else: + # otherwise, we need to convert it to target dtype + text_encoders[1].to(weight_dtype) + with accelerator.autocast(): dataset.new_cache_text_encoder_outputs(text_encoders, accelerator.is_main_process) # cache sample prompts - if args.sample_prompts is not None: logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}") @@ -210,8 +267,10 @@ class FluxNetworkTrainer(NetworkTrainer): tokenize_strategy, text_encoders, tokens_and_masks, args.apply_t5_attn_mask ) self.sample_prompts_te_outputs = sample_prompts_te_outputs + accelerator.wait_for_everyone() + # move back to cpu if not self.is_train_text_encoder(args): logger.info("move CLIP-L back to cpu") text_encoders[0].to("cpu") @@ -226,7 +285,7 @@ class FluxNetworkTrainer(NetworkTrainer): else: # Text Encoder text_encoders[0].to(accelerator.device, dtype=weight_dtype) - text_encoders[1].to(accelerator.device, dtype=weight_dtype) + text_encoders[1].to(accelerator.device) def sample_images_split_mode(self, accelerator, args, epoch, global_step, flux, ae, text_encoder, sample_prompts_te_outputs, validation_settings): @@ -259,9 +318,6 @@ class FluxNetworkTrainer(NetworkTrainer): noise_scheduler = sd3_train_utils.FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=args.discrete_flow_shift) self.noise_scheduler_copy = copy.deepcopy(noise_scheduler) return noise_scheduler - - def is_text_encoder_not_needed_for_training(self, args): - return args.cache_text_encoder_outputs and not self.is_train_text_encoder(args) def encode_images_to_latents(self, args, accelerator, vae, images): return vae.encode(images) @@ -282,55 +338,6 @@ class FluxNetworkTrainer(NetworkTrainer): weight_dtype, train_unet, ): - # copy from sd3_train.py and modified - - def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): - sigmas = self.noise_scheduler_copy.sigmas.to(device=accelerator.device, dtype=dtype) - schedule_timesteps = self.noise_scheduler_copy.timesteps.to(accelerator.device) - timesteps = timesteps.to(accelerator.device) - step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] - - sigma = sigmas[step_indices].flatten() - while len(sigma.shape) < n_dim: - sigma = sigma.unsqueeze(-1) - return sigma - - def compute_density_for_timestep_sampling( - weighting_scheme: str, batch_size: int, logit_mean: float = None, logit_std: float = None, mode_scale: float = None - ): - """Compute the density for sampling the timesteps when doing SD3 training. - - Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528. - - SD3 paper reference: https://arxiv.org/abs/2403.03206v1. - """ - if weighting_scheme == "logit_normal": - # See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$). - u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device="cpu") - u = torch.nn.functional.sigmoid(u) - elif weighting_scheme == "mode": - u = torch.rand(size=(batch_size,), device="cpu") - u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u) - else: - u = torch.rand(size=(batch_size,), device="cpu") - return u - - def compute_loss_weighting_for_sd3(weighting_scheme: str, sigmas=None): - """Computes loss weighting scheme for SD3 training. - - Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528. - - SD3 paper reference: https://arxiv.org/abs/2403.03206v1. - """ - if weighting_scheme == "sigma_sqrt": - weighting = (sigmas**-2.0).float() - elif weighting_scheme == "cosmap": - bot = 1 - 2 * sigmas + 2 * sigmas**2 - weighting = 2 / (math.pi * bot) - else: - weighting = torch.ones_like(sigmas) - return weighting - # Sample noise that we'll add to the latents noise = torch.randn_like(latents) bsz = latents.shape[0] @@ -346,7 +353,8 @@ class FluxNetworkTrainer(NetworkTrainer): img_ids = flux_utils.prepare_img_ids(bsz, packed_latent_height, packed_latent_width).to(device=accelerator.device) # get guidance - guidance_vec = torch.full((bsz,), args.guidance_scale, device=accelerator.device) + # ensure guidance_scale in args is float + guidance_vec = torch.full((bsz,), float(args.guidance_scale), device=accelerator.device) # ensure the hidden state will require grad if args.gradient_checkpointing: @@ -438,3 +446,70 @@ class FluxNetworkTrainer(NetworkTrainer): metadata["ss_sigmoid_scale"] = args.sigmoid_scale metadata["ss_model_prediction_type"] = args.model_prediction_type metadata["ss_discrete_flow_shift"] = args.discrete_flow_shift + + def is_text_encoder_not_needed_for_training(self, args): + return args.cache_text_encoder_outputs and not self.is_train_text_encoder(args) + + def prepare_text_encoder_grad_ckpt_workaround(self, index, text_encoder): + if index == 0: # CLIP-L + return super().prepare_text_encoder_grad_ckpt_workaround(index, text_encoder) + else: # T5XXL + text_encoder.encoder.embed_tokens.requires_grad_(True) + + def prepare_text_encoder_fp8(self, index, text_encoder, te_weight_dtype, weight_dtype): + if index == 0: # CLIP-L + logger.info(f"prepare CLIP-L for fp8: set to {te_weight_dtype}, set embeddings to {weight_dtype}") + text_encoder.to(te_weight_dtype) # fp8 + text_encoder.text_model.embeddings.to(dtype=weight_dtype) + else: # T5XXL + + def prepare_fp8(text_encoder, target_dtype): + def forward_hook(module): + def forward(hidden_states): + hidden_gelu = module.act(module.wi_0(hidden_states)) + hidden_linear = module.wi_1(hidden_states) + hidden_states = hidden_gelu * hidden_linear + hidden_states = module.dropout(hidden_states) + + hidden_states = module.wo(hidden_states) + return hidden_states + + return forward + + for module in text_encoder.modules(): + if module.__class__.__name__ in ["T5LayerNorm", "Embedding"]: + # print("set", module.__class__.__name__, "to", target_dtype) + module.to(target_dtype) + if module.__class__.__name__ in ["T5DenseGatedActDense"]: + # print("set", module.__class__.__name__, "hooks") + module.forward = forward_hook(module) + + if flux_utils.get_t5xxl_actual_dtype(text_encoder) == torch.float8_e4m3fn and text_encoder.dtype == weight_dtype: + logger.info(f"T5XXL already prepared for fp8") + else: + logger.info(f"prepare T5XXL for fp8: set to {te_weight_dtype}, set embeddings to {weight_dtype}, add hooks") + text_encoder.to(te_weight_dtype) # fp8 + prepare_fp8(text_encoder, weight_dtype) + + +def setup_parser() -> argparse.ArgumentParser: + parser = setup_parser() + flux_train_utils.add_flux_train_arguments(parser) + + parser.add_argument( + "--split_mode", + action="store_true", + help="[EXPERIMENTAL] use split mode for Flux model, network arg `train_blocks=single` is required" + ) + return parser + + +if __name__ == "__main__": + parser = setup_parser() + + args = parser.parse_args() + train_util.verify_command_line_training_args(args) + args = train_util.read_config_from_file(args, parser) + + trainer = FluxNetworkTrainer() + trainer.train(args) \ No newline at end of file diff --git a/library/flux_train_utils.py b/library/flux_train_utils.py index f909e9f..de65d38 100644 --- a/library/flux_train_utils.py +++ b/library/flux_train_utils.py @@ -83,7 +83,7 @@ def sample_images( except Exception: pass - with torch.no_grad(): + with torch.no_grad(), accelerator.autocast(): image_tensor_list = [] for prompt_dict in prompts: image_tensor = sample_image_inference( @@ -180,13 +180,27 @@ def sample_image_inference( tokenize_strategy = strategy_base.TokenizeStrategy.get_strategy() encoding_strategy = strategy_base.TextEncodingStrategy.get_strategy() + text_encoder_conds = [] if sample_prompts_te_outputs and prompt in sample_prompts_te_outputs: - te_outputs = sample_prompts_te_outputs[prompt] - else: + text_encoder_conds = sample_prompts_te_outputs[prompt] + print(f"Using cached text encoder outputs for prompt: {prompt}") + if text_encoders is not None: + print(f"Encoding prompt: {prompt}") tokens_and_masks = tokenize_strategy.tokenize(prompt) - te_outputs = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, tokens_and_masks) + # strategy has apply_t5_attn_mask option + encoded_text_encoder_conds = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, tokens_and_masks) + print([x.shape if x is not None else None for x in encoded_text_encoder_conds]) - l_pooled, t5_out, txt_ids, t5_attn_mask = te_outputs + # if text_encoder_conds is not cached, use encoded_text_encoder_conds + if len(text_encoder_conds) == 0: + text_encoder_conds = encoded_text_encoder_conds + else: + # if encoded_text_encoder_conds is not None, update cached text_encoder_conds + for i in range(len(encoded_text_encoder_conds)): + if encoded_text_encoder_conds[i] is not None: + text_encoder_conds[i] = encoded_text_encoder_conds[i] + + l_pooled, t5_out, txt_ids, t5_attn_mask = text_encoder_conds # sample image weight_dtype = ae.dtype # TOFO give dtype as argument @@ -522,7 +536,7 @@ def add_flux_train_arguments(parser: argparse.ArgumentParser): parser.add_argument( "--apply_t5_attn_mask", action="store_true", - help="apply attention mask (zero embs) to T5-XXL / T5-XXLにアテンションマスク(ゼロ埋め)を適用する", + help="apply attention mask to T5-XXL encode and FLUX double blocks / T5-XXLエンコードとFLUXダブルブロックにアテンションマスクを適用する", ) parser.add_argument( "--cache_text_encoder_outputs", action="store_true", help="cache text encoder outputs / text encoderの出力をキャッシュする" @@ -571,9 +585,10 @@ def add_flux_train_arguments(parser: argparse.ArgumentParser): parser.add_argument( "--timestep_sampling", - choices=["sigma", "uniform", "sigmoid"], + choices=["sigma", "uniform", "sigmoid", "shift", "flux_shift"], default="sigma", - help="Method to sample timesteps: sigma-based, uniform random, or sigmoid of random normal. / タイムステップをサンプリングする方法:sigma、random uniform、またはrandom normalのsigmoid。", + help="Method to sample timesteps: sigma-based, uniform random, sigmoid of random normal, shift of sigmoid and FLUX.1 shifting." + " / タイムステップをサンプリングする方法:sigma、random uniform、random normalのsigmoid、sigmoidのシフト、FLUX.1のシフト。", ) parser.add_argument( "--sigmoid_scale", diff --git a/library/flux_utils.py b/library/flux_utils.py index 2b987a0..f8a4c0c 100644 --- a/library/flux_utils.py +++ b/library/flux_utils.py @@ -1,5 +1,5 @@ import json -from typing import Union +from typing import Optional, Union import einops import torch @@ -8,7 +8,7 @@ from accelerate import init_empty_weights from transformers import CLIPTextModel, CLIPConfig, T5EncoderModel, T5Config from .flux_models import Flux, AutoEncoder, configs -from .utils import setup_logging +from .utils import setup_logging, MemoryEfficientSafeOpen setup_logging() import logging @@ -18,32 +18,67 @@ logger = logging.getLogger(__name__) MODEL_VERSION_FLUX_V1 = "flux1" -def load_flow_model(name: str, ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> Flux: +# temporary copy from sd3_utils TODO refactor +def load_safetensors( + path: str, device: Union[str, torch.device], disable_mmap: bool = False, dtype: Optional[torch.dtype] = torch.float32 +): + if disable_mmap: + # return safetensors.torch.load(open(path, "rb").read()) + # use experimental loader + logger.info(f"Loading without mmap (experimental)") + state_dict = {} + with MemoryEfficientSafeOpen(path) as f: + for key in f.keys(): + state_dict[key] = f.get_tensor(key).to(device, dtype=dtype) + return state_dict + else: + try: + return load_file(path, device=device) + except: + return load_file(path) # prevent device invalid Error + + +def load_flow_model( + name: str, ckpt_path: str, dtype: Optional[torch.dtype], device: Union[str, torch.device], disable_mmap: bool = False +) -> Flux: logger.info(f"Building Flux model {name}") with torch.device("meta"): - model = Flux(configs[name].params).to(dtype) + model = Flux(configs[name].params) + if dtype is not None: + model = model.to(dtype) # load_sft doesn't support torch.device logger.info(f"Loading state dict from {ckpt_path}") - sd = load_file(ckpt_path, device=str(device)) + sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype) + + # Check the first key to see if it contains the prefix + first_key = next(iter(sd)) + if first_key.startswith("model.diffusion_model."): + # Remove the 'model.diffusion_model.' prefix from keys if it exists + sd = { + key.replace("model.diffusion_model.", ""): value + for key, value in sd.items() + } info = model.load_state_dict(sd, strict=False, assign=True) logger.info(f"Loaded Flux: {info}") return model -def load_ae(name: str, ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> AutoEncoder: +def load_ae( + name: str, ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device], disable_mmap: bool = False +) -> AutoEncoder: logger.info("Building AutoEncoder") with torch.device("meta"): ae = AutoEncoder(configs[name].ae_params).to(dtype) logger.info(f"Loading state dict from {ckpt_path}") - sd = load_file(ckpt_path, device=str(device)) + sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype) info = ae.load_state_dict(sd, strict=False, assign=True) logger.info(f"Loaded AE: {info}") return ae -def load_clip_l(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> CLIPTextModel: +def load_clip_l(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device], disable_mmap: bool = False) -> CLIPTextModel: logger.info("Building CLIP") CLIPL_CONFIG = { "_name_or_path": "clip-vit-large-patch14/", @@ -138,13 +173,15 @@ def load_clip_l(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.dev clip = CLIPTextModel._from_config(config) logger.info(f"Loading state dict from {ckpt_path}") - sd = load_file(ckpt_path, device=str(device)) + sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype) info = clip.load_state_dict(sd, strict=False, assign=True) logger.info(f"Loaded CLIP: {info}") return clip -def load_t5xxl(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> T5EncoderModel: +def load_t5xxl( + ckpt_path: str, dtype: Optional[torch.dtype], device: Union[str, torch.device], disable_mmap: bool = False +) -> T5EncoderModel: T5_CONFIG_JSON = """ { "architectures": [ @@ -184,12 +221,17 @@ def load_t5xxl(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.devi t5xxl = T5EncoderModel._from_config(config) logger.info(f"Loading state dict from {ckpt_path}") - sd = load_file(ckpt_path, device=str(device)) + sd = load_safetensors(ckpt_path, device=str(device), disable_mmap=disable_mmap, dtype=dtype) info = t5xxl.load_state_dict(sd, strict=False, assign=True) logger.info(f"Loaded T5xxl: {info}") return t5xxl +def get_t5xxl_actual_dtype(t5xxl: T5EncoderModel) -> torch.dtype: + # nn.Embedding is the first layer, but it could be casted to bfloat16 or float32 + return t5xxl.encoder.block[0].layer[0].SelfAttention.q.weight.dtype + + def prepare_img_ids(batch_size: int, packed_latent_height: int, packed_latent_width: int): img_ids = torch.zeros(packed_latent_height, packed_latent_width, 3) img_ids[..., 1] = img_ids[..., 1] + torch.arange(packed_latent_height)[:, None] diff --git a/library/model_util.py b/library/model_util.py index d8d79f8..47c5674 100644 --- a/library/model_util.py +++ b/library/model_util.py @@ -1004,6 +1004,15 @@ def load_models_from_stable_diffusion_checkpoint(v2, ckpt_path, device="cpu", dt unet_config = create_unet_diffusers_config(v2, unet_use_linear_projection_in_v2) converted_unet_checkpoint = convert_ldm_unet_checkpoint(v2, state_dict, unet_config) + # convert keys of comfy saved models + first_key = next(iter(converted_unet_checkpoint)) + if first_key.startswith("model.diffusion_model."): + # Remove the 'model.diffusion_model.' prefix from keys if it exists + converted_unet_checkpoint = { + key.replace("model.diffusion_model.", ""): value + for key, value in converted_unet_checkpoint.items() + } + unet = UNet2DConditionModel(**unet_config).to(device) info = unet.load_state_dict(converted_unet_checkpoint) logger.info(f"loading u-net: {info}") diff --git a/library/strategy_flux.py b/library/strategy_flux.py index cbcb98f..65d0bc5 100644 --- a/library/strategy_flux.py +++ b/library/strategy_flux.py @@ -6,6 +6,7 @@ import numpy as np from transformers import CLIPTokenizer, T5TokenizerFast from . import train_util +from .flux_utils import get_t5xxl_actual_dtype from .strategy_base import LatentsCachingStrategy, TextEncodingStrategy, TokenizeStrategy, TextEncoderOutputsCachingStrategy from .utils import setup_logging @@ -80,7 +81,7 @@ class FluxTextEncodingStrategy(TextEncodingStrategy): else: t5_out = None txt_ids = None - t5_attn_mask = None # caption may be dropped/shuffled, so t5_attn_mask should not be used to make sure the mask is same as the cached one + t5_attn_mask = None # caption may be dropped/shuffled, so t5_attn_mask should not be used to make sure the mask is same as the cached one return [l_pooled, t5_out, txt_ids, t5_attn_mask] # returns t5_attn_mask for attention mask in transformer @@ -99,6 +100,8 @@ class FluxTextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy): super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check, is_partial) self.apply_t5_attn_mask = apply_t5_attn_mask + self.warn_fp8_weights = False + def get_outputs_npz_path(self, image_abs_path: str) -> str: return os.path.splitext(image_abs_path)[0] + FluxTextEncoderOutputsCachingStrategy.FLUX_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX @@ -144,6 +147,13 @@ class FluxTextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy): def cache_batch_outputs( self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, infos: List ): + if not self.warn_fp8_weights: + if get_t5xxl_actual_dtype(models[1]) == torch.float8_e4m3fn: + logger.warning( + "T5 model is using fp8 weights for caching. This may affect the quality of the cached outputs." + ) + self.warn_fp8_weights = True + flux_text_encoding_strategy: FluxTextEncodingStrategy = text_encoding_strategy captions = [info.caption for info in infos] diff --git a/library/train_util.py b/library/train_util.py index 9999e90..a960522 100644 --- a/library/train_util.py +++ b/library/train_util.py @@ -3521,6 +3521,13 @@ def add_training_arguments(parser: argparse.ArgumentParser, support_dreambooth: "--full_bf16", action="store_true", help="bf16 training including gradients / 勾配も含めてbf16で学習する" ) # TODO move to SDXL training, because it is not supported by SD1/2 parser.add_argument("--fp8_base", action="store_true", help="use fp8 for base model / base modelにfp8を使う") + parser.add_argument( + "--fp8_dtype", + type=str, + default="e4m3", + choices=["e4m3", "e5m2"], + help="fp8 dtype selection", + ) parser.add_argument( "--ddp_timeout", @@ -4782,6 +4789,8 @@ def prepare_dtype(args: argparse.Namespace): save_dtype = torch.float32 elif args.save_precision == "fp8_e4m3fn": save_dtype = torch.float8_e4m3fn + elif args.save_precision == "fp8_e5m2": + save_dtype = torch.float8_e5m2 return weight_dtype, save_dtype diff --git a/networks/lora_flux.py b/networks/lora_flux.py index a7c99d8..0878346 100644 --- a/networks/lora_flux.py +++ b/networks/lora_flux.py @@ -38,6 +38,7 @@ class LoRAModule(torch.nn.Module): dropout=None, rank_dropout=None, module_dropout=None, + split_dims: Optional[List[int]] = None, ): """if alpha == 0 or None, alpha is rank (no scaling).""" super().__init__() @@ -51,16 +52,34 @@ class LoRAModule(torch.nn.Module): out_dim = org_module.out_features self.lora_dim = lora_dim + self.split_dims = split_dims - if org_module.__class__.__name__ == "Conv2d": - kernel_size = org_module.kernel_size - stride = org_module.stride - padding = org_module.padding - self.lora_down = torch.nn.Conv2d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False) - self.lora_up = torch.nn.Conv2d(self.lora_dim, out_dim, (1, 1), (1, 1), bias=False) + if split_dims is None: + if org_module.__class__.__name__ == "Conv2d": + kernel_size = org_module.kernel_size + stride = org_module.stride + padding = org_module.padding + self.lora_down = torch.nn.Conv2d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False) + self.lora_up = torch.nn.Conv2d(self.lora_dim, out_dim, (1, 1), (1, 1), bias=False) + else: + self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False) + self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=False) + + torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5)) + torch.nn.init.zeros_(self.lora_up.weight) else: - self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False) - self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=False) + # conv2d not supported + assert sum(split_dims) == out_dim, "sum of split_dims must be equal to out_dim" + assert org_module.__class__.__name__ == "Linear", "split_dims is only supported for Linear" + # print(f"split_dims: {split_dims}") + self.lora_down = torch.nn.ModuleList( + [torch.nn.Linear(in_dim, self.lora_dim, bias=False) for _ in range(len(split_dims))] + ) + self.lora_up = torch.nn.ModuleList([torch.nn.Linear(self.lora_dim, split_dim, bias=False) for split_dim in split_dims]) + for lora_down in self.lora_down: + torch.nn.init.kaiming_uniform_(lora_down.weight, a=math.sqrt(5)) + for lora_up in self.lora_up: + torch.nn.init.zeros_(lora_up.weight) if type(alpha) == torch.Tensor: alpha = alpha.detach().float().numpy() # without casting, bf16 causes error @@ -69,9 +88,6 @@ class LoRAModule(torch.nn.Module): self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える # same as microsoft's - torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5)) - torch.nn.init.zeros_(self.lora_up.weight) - self.multiplier = multiplier self.org_module = org_module # remove in applying self.dropout = dropout @@ -91,30 +107,56 @@ class LoRAModule(torch.nn.Module): if torch.rand(1) < self.module_dropout: return org_forwarded - lx = self.lora_down(x) + if self.split_dims is None: + lx = self.lora_down(x) - # normal dropout - if self.dropout is not None and self.training: - lx = torch.nn.functional.dropout(lx, p=self.dropout) + # normal dropout + if self.dropout is not None and self.training: + lx = torch.nn.functional.dropout(lx, p=self.dropout) - # rank dropout - if self.rank_dropout is not None and self.training: - mask = torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout - if len(lx.size()) == 3: - mask = mask.unsqueeze(1) # for Text Encoder - elif len(lx.size()) == 4: - mask = mask.unsqueeze(-1).unsqueeze(-1) # for Conv2d - lx = lx * mask + # rank dropout + if self.rank_dropout is not None and self.training: + mask = torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout + if len(lx.size()) == 3: + mask = mask.unsqueeze(1) # for Text Encoder + elif len(lx.size()) == 4: + mask = mask.unsqueeze(-1).unsqueeze(-1) # for Conv2d + lx = lx * mask - # scaling for rank dropout: treat as if the rank is changed - # maskから計算することも考えられるが、augmentation的な効果を期待してrank_dropoutを用いる - scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability + # scaling for rank dropout: treat as if the rank is changed + # maskから計算することも考えられるが、augmentation的な効果を期待してrank_dropoutを用いる + scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability + else: + scale = self.scale + + lx = self.lora_up(lx) + + return org_forwarded + lx * self.multiplier * scale else: - scale = self.scale + lxs = [lora_down(x) for lora_down in self.lora_down] - lx = self.lora_up(lx) + # normal dropout + if self.dropout is not None and self.training: + lxs = [torch.nn.functional.dropout(lx, p=self.dropout) for lx in lxs] - return org_forwarded + lx * self.multiplier * scale + # rank dropout + if self.rank_dropout is not None and self.training: + masks = [torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout for lx in lxs] + for i in range(len(lxs)): + if len(lx.size()) == 3: + masks[i] = masks[i].unsqueeze(1) + elif len(lx.size()) == 4: + masks[i] = masks[i].unsqueeze(-1).unsqueeze(-1) + lxs[i] = lxs[i] * masks[i] + + # scaling for rank dropout: treat as if the rank is changed + scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability + else: + scale = self.scale + + lxs = [lora_up(lx) for lora_up, lx in zip(self.lora_up, lxs)] + + return org_forwarded + torch.cat(lxs, dim=-1) * self.multiplier * scale class LoRAInfModule(LoRAModule): @@ -151,31 +193,50 @@ class LoRAInfModule(LoRAModule): if device is None: device = org_device - # get up/down weight - up_weight = sd["lora_up.weight"].to(torch.float).to(device) - down_weight = sd["lora_down.weight"].to(torch.float).to(device) + if self.split_dims is None: + # get up/down weight + down_weight = sd["lora_down.weight"].to(torch.float).to(device) + up_weight = sd["lora_up.weight"].to(torch.float).to(device) - # merge weight - if len(weight.size()) == 2: - # linear - weight = weight + self.multiplier * (up_weight @ down_weight) * self.scale - elif down_weight.size()[2:4] == (1, 1): - # conv2d 1x1 - weight = ( - weight - + self.multiplier - * (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3) - * self.scale - ) + # merge weight + if len(weight.size()) == 2: + # linear + weight = weight + self.multiplier * (up_weight @ down_weight) * self.scale + elif down_weight.size()[2:4] == (1, 1): + # conv2d 1x1 + weight = ( + weight + + self.multiplier + * (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3) + * self.scale + ) + else: + # conv2d 3x3 + conved = torch.nn.functional.conv2d(down_weight.permute(1, 0, 2, 3), up_weight).permute(1, 0, 2, 3) + # logger.info(conved.size(), weight.size(), module.stride, module.padding) + weight = weight + self.multiplier * conved * self.scale + + # set weight to org_module + org_sd["weight"] = weight.to(dtype) + self.org_module.load_state_dict(org_sd) else: - # conv2d 3x3 - conved = torch.nn.functional.conv2d(down_weight.permute(1, 0, 2, 3), up_weight).permute(1, 0, 2, 3) - # logger.info(conved.size(), weight.size(), module.stride, module.padding) - weight = weight + self.multiplier * conved * self.scale + # split_dims + total_dims = sum(self.split_dims) + for i in range(len(self.split_dims)): + # get up/down weight + down_weight = sd[f"lora_down.{i}.weight"].to(torch.float).to(device) # (rank, in_dim) + up_weight = sd[f"lora_up.{i}.weight"].to(torch.float).to(device) # (split dim, rank) - # set weight to org_module - org_sd["weight"] = weight.to(dtype) - self.org_module.load_state_dict(org_sd) + # pad up_weight -> (total_dims, rank) + padded_up_weight = torch.zeros((total_dims, up_weight.size(0)), device=device, dtype=torch.float) + padded_up_weight[sum(self.split_dims[:i]) : sum(self.split_dims[: i + 1])] = up_weight + + # merge weight + weight = weight + self.multiplier * (up_weight @ down_weight) * self.scale + + # set weight to org_module + org_sd["weight"] = weight.to(dtype) + self.org_module.load_state_dict(org_sd) # 復元できるマージのため、このモジュールのweightを返す def get_weight(self, multiplier=None): @@ -210,7 +271,14 @@ class LoRAInfModule(LoRAModule): def default_forward(self, x): # logger.info(f"default_forward {self.lora_name} {x.size()}") - return self.org_forward(x) + self.lora_up(self.lora_down(x)) * self.multiplier * self.scale + if self.split_dims is None: + lx = self.lora_down(x) + lx = self.lora_up(lx) + return self.org_forward(x) + lx * self.multiplier * self.scale + else: + lxs = [lora_down(x) for lora_down in self.lora_down] + lxs = [lora_up(lx) for lora_up, lx in zip(self.lora_up, lxs)] + return self.org_forward(x) + torch.cat(lxs, dim=-1) * self.multiplier * self.scale def forward(self, x): if not self.enabled: @@ -256,6 +324,20 @@ def create_network( if train_blocks is not None: assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}" + only_if_contains = kwargs.get("only_if_contains", None) + if only_if_contains is not None: + only_if_contains = [word.strip() for word in only_if_contains.split(',')] + + # split qkv + split_qkv = kwargs.get("split_qkv", False) + if split_qkv is not None: + split_qkv = True if split_qkv == "True" else False + + # train T5XXL + train_t5xxl = kwargs.get("train_t5xxl", False) + if train_t5xxl is not None: + train_t5xxl = True if train_t5xxl == "True" else False + # すごく引数が多いな ( ^ω^)・・・ network = LoRANetwork( text_encoders, @@ -269,7 +351,10 @@ def create_network( conv_lora_dim=conv_dim, conv_alpha=conv_alpha, train_blocks=train_blocks, + split_qkv=split_qkv, + train_t5xxl=train_t5xxl, varbose=True, + only_if_contains=only_if_contains ) loraplus_lr_ratio = kwargs.get("loraplus_lr_ratio", None) @@ -295,9 +380,10 @@ def create_network_from_weights(multiplier, file, ae, text_encoders, flux, weigh else: weights_sd = torch.load(file, map_location="cpu") - # get dim/alpha mapping + # get dim/alpha mapping, and train t5xxl modules_dim = {} modules_alpha = {} + train_t5xxl = None for key, value in weights_sd.items(): if "." not in key: continue @@ -310,10 +396,41 @@ def create_network_from_weights(multiplier, file, ae, text_encoders, flux, weigh modules_dim[lora_name] = dim # logger.info(lora_name, value.size(), dim) + if train_t5xxl is None or train_t5xxl is False: + train_t5xxl = "lora_te3" in lora_name + + if train_t5xxl is None: + train_t5xxl = False + + # # split qkv + # double_qkv_rank = None + # single_qkv_rank = None + # rank = None + # for lora_name, dim in modules_dim.items(): + # if "double" in lora_name and "qkv" in lora_name: + # double_qkv_rank = dim + # elif "single" in lora_name and "linear1" in lora_name: + # single_qkv_rank = dim + # elif rank is None: + # rank = dim + # if double_qkv_rank is not None and single_qkv_rank is not None and rank is not None: + # break + # split_qkv = (double_qkv_rank is not None and double_qkv_rank != rank) or ( + # single_qkv_rank is not None and single_qkv_rank != rank + # ) + split_qkv = False # split_qkv is not needed to care, because state_dict is qkv combined + module_class = LoRAInfModule if for_inference else LoRAModule network = LoRANetwork( - text_encoders, flux, multiplier=multiplier, modules_dim=modules_dim, modules_alpha=modules_alpha, module_class=module_class + text_encoders, + flux, + multiplier=multiplier, + modules_dim=modules_dim, + modules_alpha=modules_alpha, + module_class=module_class, + split_qkv=split_qkv, + train_t5xxl=train_t5xxl, ) return network, weights_sd @@ -322,10 +439,10 @@ class LoRANetwork(torch.nn.Module): # FLUX_TARGET_REPLACE_MODULE = ["DoubleStreamBlock", "SingleStreamBlock"] FLUX_TARGET_REPLACE_MODULE_DOUBLE = ["DoubleStreamBlock"] FLUX_TARGET_REPLACE_MODULE_SINGLE = ["SingleStreamBlock"] - TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPSdpaAttention", "CLIPMLP"] + TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPSdpaAttention", "CLIPMLP", "T5Attention", "T5DenseGatedActDense"] LORA_PREFIX_FLUX = "lora_unet" # make ComfyUI compatible LORA_PREFIX_TEXT_ENCODER_CLIP = "lora_te1" - LORA_PREFIX_TEXT_ENCODER_T5 = "lora_te2" + LORA_PREFIX_TEXT_ENCODER_T5 = "lora_te3" # make ComfyUI compatible def __init__( self, @@ -343,7 +460,10 @@ class LoRANetwork(torch.nn.Module): modules_dim: Optional[Dict[str, int]] = None, modules_alpha: Optional[Dict[str, int]] = None, train_blocks: Optional[str] = None, + split_qkv: bool = False, + train_t5xxl: bool = False, varbose: Optional[bool] = False, + only_if_contains: Optional[List[str]] = None, ) -> None: super().__init__() self.multiplier = multiplier @@ -356,11 +476,15 @@ class LoRANetwork(torch.nn.Module): self.rank_dropout = rank_dropout self.module_dropout = module_dropout self.train_blocks = train_blocks if train_blocks is not None else "all" + self.split_qkv = split_qkv + self.train_t5xxl = train_t5xxl self.loraplus_lr_ratio = None self.loraplus_unet_lr_ratio = None self.loraplus_text_encoder_lr_ratio = None + self.only_if_contains = only_if_contains + if modules_dim is not None: logger.info(f"create LoRA network from weights") else: @@ -368,10 +492,18 @@ class LoRANetwork(torch.nn.Module): logger.info( f"neuron dropout: p={self.dropout}, rank dropout: p={self.rank_dropout}, module dropout: p={self.module_dropout}" ) - if self.conv_lora_dim is not None: - logger.info( - f"apply LoRA to Conv2d with kernel size (3,3). dim (rank): {self.conv_lora_dim}, alpha: {self.conv_alpha}" - ) + # if self.conv_lora_dim is not None: + # logger.info( + # f"apply LoRA to Conv2d with kernel size (3,3). dim (rank): {self.conv_lora_dim}, alpha: {self.conv_alpha}" + # ) + if self.split_qkv: + logger.info(f"split qkv for LoRA") + if self.train_blocks is not None: + logger.info(f"train {self.train_blocks} blocks only") + if train_t5xxl: + logger.info(f"train T5XXL as well") + + #self.only_if_contains = ["lora_unet_single_blocks_20_linear2"] # create module instances def create_modules( @@ -395,6 +527,10 @@ class LoRANetwork(torch.nn.Module): if is_linear or is_conv2d: lora_name = prefix + "." + name + "." + child_name lora_name = lora_name.replace(".", "_") + #lora_unet_single_blocks_20_linear2 + + if "unet" in lora_name and (self.only_if_contains is not None and not any(word in lora_name for word in self.only_if_contains)): + continue dim = None alpha = None @@ -419,6 +555,14 @@ class LoRANetwork(torch.nn.Module): skipped.append(lora_name) continue + # qkv split + split_dims = None + if is_flux and split_qkv: + if "double" in lora_name and "qkv" in lora_name: + split_dims = [3072] * 3 + elif "single" in lora_name and "linear1" in lora_name: + split_dims = [3072] * 3 + [12288] + lora = module_class( lora_name, child_module, @@ -428,6 +572,7 @@ class LoRANetwork(torch.nn.Module): dropout=dropout, rank_dropout=rank_dropout, module_dropout=module_dropout, + split_dims=split_dims, ) loras.append(lora) return loras, skipped @@ -438,12 +583,15 @@ class LoRANetwork(torch.nn.Module): skipped_te = [] for i, text_encoder in enumerate(text_encoders): index = i + if not train_t5xxl and index > 0: # 0: CLIP, 1: T5XXL, so we skip T5XXL if train_t5xxl is False + break + logger.info(f"create LoRA for Text Encoder {index+1}:") text_encoder_loras, skipped = create_modules(False, index, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE) + logger.info(f"create LoRA for Text Encoder {index+1}: {len(text_encoder_loras)} modules.") self.text_encoder_loras.extend(text_encoder_loras) skipped_te += skipped - logger.info(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.") # create LoRA for U-Net if self.train_blocks == "all": @@ -456,6 +604,7 @@ class LoRANetwork(torch.nn.Module): self.unet_loras: List[Union[LoRAModule, LoRAInfModule]] self.unet_loras, skipped_un = create_modules(True, None, unet, target_replace_modules) logger.info(f"create LoRA for FLUX {self.train_blocks} blocks: {len(self.unet_loras)} modules.") + print(self.unet_loras) skipped = skipped_te + skipped_un if varbose and len(skipped) > 0: @@ -491,6 +640,111 @@ class LoRANetwork(torch.nn.Module): info = self.load_state_dict(weights_sd, False) return info + def load_state_dict(self, state_dict, strict=True): + # override to convert original weight to split qkv + if not self.split_qkv: + return super().load_state_dict(state_dict, strict) + + # split qkv + for key in list(state_dict.keys()): + if "double" in key and "qkv" in key: + split_dims = [3072] * 3 + elif "single" in key and "linear1" in key: + split_dims = [3072] * 3 + [12288] + else: + continue + + weight = state_dict[key] + lora_name = key.split(".")[0] + if "lora_down" in key and "weight" in key: + # dense weight (rank*3, in_dim) + split_weight = torch.chunk(weight, len(split_dims), dim=0) + for i, split_w in enumerate(split_weight): + state_dict[f"{lora_name}.lora_down.{i}.weight"] = split_w + + del state_dict[key] + # print(f"split {key}: {weight.shape} to {[w.shape for w in split_weight]}") + elif "lora_up" in key and "weight" in key: + # sparse weight (out_dim=sum(split_dims), rank*3) + rank = weight.size(1) // len(split_dims) + i = 0 + for j in range(len(split_dims)): + state_dict[f"{lora_name}.lora_up.{j}.weight"] = weight[i : i + split_dims[j], j * rank : (j + 1) * rank] + i += split_dims[j] + del state_dict[key] + + # # check is sparse + # i = 0 + # is_zero = True + # for j in range(len(split_dims)): + # for k in range(len(split_dims)): + # if j == k: + # continue + # is_zero = is_zero and torch.all(weight[i : i + split_dims[j], k * rank : (k + 1) * rank] == 0) + # i += split_dims[j] + # if not is_zero: + # logger.warning(f"weight is not sparse: {key}") + # else: + # logger.info(f"weight is sparse: {key}") + + # print( + # f"split {key}: {weight.shape} to {[state_dict[k].shape for k in [f'{lora_name}.lora_up.{j}.weight' for j in range(len(split_dims))]]}" + # ) + + # alpha is unchanged + + return super().load_state_dict(state_dict, strict) + + def state_dict(self, destination=None, prefix="", keep_vars=False): + if not self.split_qkv: + return super().state_dict(destination, prefix, keep_vars) + + # merge qkv + state_dict = super().state_dict(destination, prefix, keep_vars) + new_state_dict = {} + for key in list(state_dict.keys()): + if "double" in key and "qkv" in key: + split_dims = [3072] * 3 + elif "single" in key and "linear1" in key: + split_dims = [3072] * 3 + [12288] + else: + new_state_dict[key] = state_dict[key] + continue + + if key not in state_dict: + continue # already merged + + lora_name = key.split(".")[0] + + # (rank, in_dim) * 3 + down_weights = [state_dict.pop(f"{lora_name}.lora_down.{i}.weight") for i in range(len(split_dims))] + # (split dim, rank) * 3 + up_weights = [state_dict.pop(f"{lora_name}.lora_up.{i}.weight") for i in range(len(split_dims))] + + alpha = state_dict.pop(f"{lora_name}.alpha") + + # merge down weight + down_weight = torch.cat(down_weights, dim=0) # (rank, split_dim) * 3 -> (rank*3, sum of split_dim) + + # merge up weight (sum of split_dim, rank*3) + rank = up_weights[0].size(1) + up_weight = torch.zeros((sum(split_dims), down_weight.size(0)), device=down_weight.device, dtype=down_weight.dtype) + i = 0 + for j in range(len(split_dims)): + up_weight[i : i + split_dims[j], j * rank : (j + 1) * rank] = up_weights[j] + i += split_dims[j] + + new_state_dict[f"{lora_name}.lora_down.weight"] = down_weight + new_state_dict[f"{lora_name}.lora_up.weight"] = up_weight + new_state_dict[f"{lora_name}.alpha"] = alpha + + # print( + # f"merged {lora_name}: {lora_name}, {[w.shape for w in down_weights]}, {[w.shape for w in up_weights]} to {down_weight.shape}, {up_weight.shape}" + # ) + print(f"new key: {lora_name}.lora_down.weight, {lora_name}.lora_up.weight, {lora_name}.alpha") + + return new_state_dict + def apply_to(self, text_encoders, flux, apply_text_encoder=True, apply_unet=True): if apply_text_encoder: logger.info(f"enable LoRA for text encoder: {len(self.text_encoder_loras)} modules") diff --git a/nodes.py b/nodes.py index 1fb4cb1..23d75d9 100644 --- a/nodes.py +++ b/nodes.py @@ -9,6 +9,8 @@ import toml import json import time import shutil +import shlex + from pathlib import Path script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -35,11 +37,14 @@ class FluxTrainModelSelect: @classmethod def INPUT_TYPES(s): return {"required": { - "transformer": (folder_paths.get_filename_list("unet"), ), - "vae": (folder_paths.get_filename_list("vae"), ), - "clip_l": (folder_paths.get_filename_list("clip"), ), - "t5": (folder_paths.get_filename_list("clip"), ), - }, + "transformer": (folder_paths.get_filename_list("unet"), ), + "vae": (folder_paths.get_filename_list("vae"), ), + "clip_l": (folder_paths.get_filename_list("clip"), ), + "t5": (folder_paths.get_filename_list("clip"), ), + }, + "optional": { + "lora_path": ("STRING",{"multiline": True, "forceInput": True, "default": "", "tooltip": "pre-trained LoRA path to load (network_weights)"}), + } } RETURN_TYPES = ("TRAIN_FLUX_MODELS",) @@ -47,7 +52,7 @@ class FluxTrainModelSelect: FUNCTION = "loadmodel" CATEGORY = "FluxTrainer" - def loadmodel(self, transformer, vae, clip_l, t5): + def loadmodel(self, transformer, vae, clip_l, t5, lora_path=""): transformer_path = folder_paths.get_full_path("unet", transformer) vae_path = folder_paths.get_full_path("vae", vae) @@ -58,12 +63,20 @@ class FluxTrainModelSelect: "transformer": transformer_path, "vae": vae_path, "clip_l": clip_path, - "t5": t5_path + "t5": t5_path, + "lora_path": lora_path } return (flux_models,) class TrainDatasetGeneralConfig: + queue_counter = 0 + @classmethod + def IS_CHANGED(s, reset_on_queue=False, **kwargs): + if reset_on_queue: + s.queue_counter += 1 + print(f"queue_counter: {s.queue_counter}") + return s.queue_counter @classmethod def INPUT_TYPES(s): return {"required": { @@ -73,6 +86,10 @@ class TrainDatasetGeneralConfig: "caption_dropout_rate": ("FLOAT",{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01,"tooltip": "tag dropout rate"}), "alpha_mask": ("BOOLEAN",{"default": False, "tooltip": "use alpha channel as mask for training"}), }, + "optional": { + "reset_on_queue": ("BOOLEAN",{"default": False, "tooltip": "Force refresh of everything for cleaner queueing"}), + "reg_data_dir": ("STRING",{"multiline": True, "forceInput": True, "default": "", "tooltip": "reg data dir"}), + } } RETURN_TYPES = ("JSON",) @@ -80,7 +97,7 @@ class TrainDatasetGeneralConfig: FUNCTION = "create_config" CATEGORY = "FluxTrainer" - def create_config(self, shuffle_caption, caption_dropout_rate, color_aug, flip_aug, alpha_mask): + def create_config(self, shuffle_caption, caption_dropout_rate, color_aug, flip_aug, alpha_mask, reset_on_queue=False, reg_data_dir=""): dataset = { "general": { @@ -97,13 +114,15 @@ class TrainDatasetGeneralConfig: #print(dataset_json) dataset_config = { "datasets": dataset_json, - "alpha_mask": alpha_mask + "alpha_mask": alpha_mask, + "reg_data_dir": reg_data_dir } return (dataset_config,) class TrainDatasetAdd: def __init__(self): self.previous_dataset_signature = None + @classmethod def INPUT_TYPES(s): return {"required": { @@ -118,7 +137,6 @@ class TrainDatasetAdd: "num_repeats": ("INT", {"default": 1, "min": 1, "tooltip": "number of times to repeat dataset for an epoch"}), "min_bucket_reso": ("INT", {"default": 256, "min": 64, "max": 4096, "step": 8, "tooltip": "min bucket resolution"}), "max_bucket_reso": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 8, "tooltip": "max bucket resolution"}), - }, } @@ -197,7 +215,7 @@ class OptimizerConfig: def create_config(self, min_snr_gamma, extra_optimizer_args, **kwargs): kwargs["min_snr_gamma"] = min_snr_gamma if min_snr_gamma != 0.0 else None - kwargs["optimizer_args"] = [arg.strip() for arg in extra_optimizer_args.strip().split(',') if arg.strip()] + kwargs["optimizer_args"] = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()] return (kwargs,) class OptimizerConfigAdafactor: @@ -225,7 +243,7 @@ class OptimizerConfigAdafactor: def create_config(self, relative_step, scale_parameter, warmup_init, clip_threshold, min_snr_gamma, extra_optimizer_args, **kwargs): kwargs["optimizer_type"] = "adafactor" - extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split(',') if arg.strip()] + extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()] node_args = [ f"relative_step={relative_step}", f"scale_parameter={scale_parameter}", @@ -261,7 +279,7 @@ class OptimizerConfigProdigy: def create_config(self, weight_decay, decouple, min_snr_gamma, use_bias_correction, extra_optimizer_args, **kwargs): kwargs["optimizer_type"] = "prodigy" - extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split(',') if arg.strip()] + extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()] node_args = [ f"weight_decay={weight_decay}", f"decouple={decouple}", @@ -284,12 +302,8 @@ class InitFluxLoRATraining: "network_dim": ("INT", {"default": 4, "min": 1, "max": 256, "step": 1, "tooltip": "network dim"}), "network_alpha": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 256.0, "step": 0.01, "tooltip": "network alpha"}), "learning_rate": ("FLOAT", {"default": 4e-4, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "learning rate"}), - #"unet_lr": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "unet learning rate"}), - #"max_train_epochs": ("INT", {"default": 4, "min": 1, "max": 1000, "step": 1, "tooltip": "max number of training epochs"}), "max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 100000, "step": 1, "tooltip": "max number of training steps"}), - #"text_encoder_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}), "apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "apply t5 attention mask"}), - #"t5xxl_max_token_length": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "dev uses 512, schnell 256"}), "cache_latents": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}), "cache_text_encoder_outputs": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}), "split_mode": ("BOOLEAN", {"default": False, "tooltip": "[EXPERIMENTAL] use split mode for Flux model, network arg `train_blocks=single` is required"}), @@ -305,15 +319,17 @@ class InitFluxLoRATraining: "highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}), "fp8_base": ("BOOLEAN", {"default": True, "tooltip": "use fp8 for base model"}), "gradient_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "the actual dtype training uses"}), - "save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}), + "save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}), "attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}), "sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}), }, "optional": { "additional_args": ("STRING", {"multiline": True, "default": "", "tooltip": "additional args to pass to the training command"}), "resume_args": ("ARGS", {"default": "", "tooltip": "resume args to pass to the training command"}), - "train_clip_l": (['disabled', 'use_gradient_dtype', 'use_fp8'], {"default": 'disabled', "tooltip": "also train the clip_l text encoder using specified dtype"}), - "text_encoder_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}), + "train_text_encoder": (['disabled', 'clip_l', 'clip_l_fp8', 'clip_l+T5', 'clip_l+T5_fp8'], {"default": 'disabled', "tooltip": "also train the selected text encoders using specified dtype, T5 can not be trained without clip_l"}), + "text_encoder_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "text encoder learning rate"}), + "block_args": ("ARGS", {"default": "", "tooltip": "limit the blocks used in the LoRA"}), + "gradient_checkpointing": (["enabled", "enabled_with_cpu_offloading", "disabled"], {"default": "enabled", "tooltip": "use gradient checkpointing"}), }, } @@ -323,7 +339,8 @@ class InitFluxLoRATraining: CATEGORY = "FluxTrainer" def init_training(self, flux_models, dataset, optimizer_settings, sample_prompts, output_name, attention_mode, - gradient_dtype, save_dtype, split_mode, additional_args=None, resume_args=None, train_clip_l='disabled', **kwargs,): + gradient_dtype, save_dtype, split_mode, additional_args=None, resume_args=None, train_text_encoder='disabled', + block_args=None, gradient_checkpointing="enabled", **kwargs,): mm.soft_empty_cache() output_dir = os.path.abspath(kwargs.get("output_dir")) @@ -340,7 +357,8 @@ class InitFluxLoRATraining: parser = train_network_setup_parser() if additional_args is not None: - args, _ = parser.parse_known_args(args=[additional_args]) + print(f"additional_args: {additional_args}") + args, _ = parser.parse_known_args(args=shlex.split(additional_args)) else: args, _ = parser.parse_known_args() #print(args) @@ -387,7 +405,6 @@ class InitFluxLoRATraining: "persistent_data_loader_workers": False, "max_data_loader_n_workers": 0, "seed": 42, - "gradient_checkpointing": True, "network_module": ".networks.lora_flux", "dataset_config": dataset_toml, "output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}", @@ -395,8 +412,10 @@ class InitFluxLoRATraining: "text_encoder_lr": 0, "t5xxl_max_token_length": 512, "alpha_mask": dataset["alpha_mask"], - "network_train_unet_only": True if train_clip_l == 'disabled' else False, - "fp8_base_unet": True if train_clip_l=='use_gradient_dtype' else False, + "network_train_unet_only": True if train_text_encoder == 'disabled' else False, + "fp8_base_unet": False if "fp8" in train_text_encoder else True, + "disable_mmap_load_safetensors": False, + "split_mode": split_mode, } attention_settings = { "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True}, @@ -410,11 +429,34 @@ class InitFluxLoRATraining: } config_dict.update(gradient_dtype_settings.get(gradient_dtype, {})) - split_mode_settings = { - True: {"split_mode": True, "network_args": ["train_blocks=single"]}, - False: {"split_mode": False, "network_args": ["train_blocks=all"]} - } - config_dict.update(split_mode_settings.get(split_mode, {})) + #network args + additional_network_args = [] + + if "T5" in train_text_encoder: + additional_network_args.append("train_t5xxl=True") + if split_mode: + additional_network_args.append("train_blocks=single") + if block_args: + additional_network_args.append(block_args["include"]) + + # Handle network_args in args Namespace + if hasattr(args, 'network_args') and isinstance(args.network_args, list): + args.network_args.extend(additional_network_args) + else: + setattr(args, 'network_args', additional_network_args) + + if gradient_checkpointing == "disabled": + config_dict["gradient_checkpointing"] = False + elif gradient_checkpointing == "enabled_with_cpu_offloading": + config_dict["cpu_offload_checkpointing"] = True + else: + config_dict["gradient_checkpointing"] = True + + if flux_models["lora_path"]: + config_dict["network_weights"] = flux_models["lora_path"] + + if dataset["reg_data_dir"]: + config_dict["reg_data_dir"] = dataset["reg_data_dir"] config_dict.update(kwargs) config_dict.update(optimizer_settings) @@ -554,6 +596,7 @@ class InitFluxTraining: "dataset_config": dataset_toml, "output_name": f"{output_name}_{save_dtype}", "mem_eff_save": True, + "disable_mmap_load_safetensors": True, } optimizer_fusing_settings = { @@ -697,18 +740,19 @@ class FluxTrainLoop: initial_global_step = network_trainer.global_step target_global_step = network_trainer.global_step + steps - pbar = comfy.utils.ProgressBar(steps) + comfy_pbar = comfy.utils.ProgressBar(steps) + network_trainer.comfy_pbar = comfy_pbar while network_trainer.global_step < target_global_step: steps_done = training_loop( break_at_steps = target_global_step, epoch = network_trainer.current_epoch.value, ) - pbar.update(steps_done) + #pbar.update(steps_done) # Also break if the global steps have reached the max train steps if network_trainer.global_step >= network_trainer.args.max_train_steps: break - + trainer = { "network_trainer": network_trainer, "training_loop": training_loop, @@ -865,6 +909,26 @@ class FluxTrainResume: return (resume_args, ) +class FluxTrainBlockSelect: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "include": ("STRING", {"default": "lora_unet_single_blocks_20_linear2", "multiline": True, "tooltip": "blocks to include in the LoRA network"}), + }, + } + + RETURN_TYPES = ("ARGS", ) + RETURN_NAMES = ("block_args", ) + FUNCTION = "block_select" + CATEGORY = "FluxTrainer" + + def block_select(self, include): + block_args ={ + "include": f"only_if_contains={include}", + } + + return (block_args, ) + class FluxTrainValidationSettings: @classmethod def INPUT_TYPES(s): @@ -1009,6 +1073,7 @@ class FluxKohyaInferenceSampler: "guidance_scale": ("FLOAT", {"default": 3.5, "min": 1.0, "max": 32.0, "step": 0.05, "tooltip": "guidance scale"}), "seed": ("INT", {"default": 42,"min": 0, "max": 0xffffffffffffffff, "step": 1}), "use_fp8": ("BOOLEAN", {"default": True, "tooltip": "use fp8 weights"}), + "apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "use t5 attention mask"}), "prompt": ("STRING", {"multiline": True, "default": "illustration of a kitten", "tooltip": "prompt"}), }, @@ -1019,7 +1084,7 @@ class FluxKohyaInferenceSampler: FUNCTION = "sample" CATEGORY = "FluxTrainer" - def sample(self, flux_models, lora_name, steps, width, height, guidance_scale, seed, prompt, use_fp8, lora_method): + def sample(self, flux_models, lora_name, steps, width, height, guidance_scale, seed, prompt, use_fp8, lora_method, apply_t5_attn_mask): from .library import flux_utils as flux_utils from .library import strategy_flux as strategy_flux @@ -1032,7 +1097,7 @@ class FluxKohyaInferenceSampler: import gc device = "cuda" - apply_t5_attn_mask = True + if use_fp8: accelerator = accelerate.Accelerator(mixed_precision="bf16") @@ -1077,8 +1142,7 @@ class FluxKohyaInferenceSampler: # AE ae = flux_utils.load_ae("dev", ae, ae_dtype, loading_device) ae.eval() - #if is_fp8(ae_dtype): - # ae = accelerator.prepare(ae) + # LoRA lora_models: List[lora_flux.LoRANetwork] = [] @@ -1120,7 +1184,7 @@ class FluxKohyaInferenceSampler: clip_l.to(ae_dtype) t5xxl.to(ae_dtype) with accelerator.autocast(): - _, t5_out, txt_ids, t5_attn_mask = encoding_strategy.encode_tokens( + l_pooled, t5_out, txt_ids, t5_attn_mask = encoding_strategy.encode_tokens( tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, apply_t5_attn_mask ) else: @@ -1226,6 +1290,7 @@ class FluxKohyaInferenceSampler: flux_dtype: torch.dtype, ): timesteps = get_schedule(num_steps, img.shape[1], shift=not is_schnell) + print(timesteps) # denoise initial noise if accelerator: @@ -1234,9 +1299,11 @@ class FluxKohyaInferenceSampler: model, img, img_ids, t5_out, txt_ids, l_pooled, timesteps=timesteps, guidance=guidance, t5_attn_mask=t5_attn_mask ) else: - with torch.autocast(device_type=device.type, dtype=flux_dtype), torch.no_grad(): - x = denoise( - model, img, img_ids, t5_out, txt_ids, l_pooled, timesteps=timesteps, guidance=guidance, t5_attn_mask=t5_attn_mask + with torch.autocast(device_type=device.type, dtype=flux_dtype): + l_pooled, _, _, _ = encoding_strategy.encode_tokens(tokenize_strategy, [clip_l, None], tokens_and_masks) + with torch.autocast(device_type=device.type, dtype=flux_dtype): + _, t5_out, txt_ids, t5_attn_mask = encoding_strategy.encode_tokens( + tokenize_strategy, [None, t5xxl], tokens_and_masks, apply_t5_attn_mask ) return x @@ -1375,7 +1442,7 @@ class ExtractFluxLoRA: "finetuned_model": (folder_paths.get_filename_list("unet"), ), "output_path": ("STRING", {"default": f"{str(os.path.join(folder_paths.models_dir, 'loras', 'Flux'))}"}), "dim": ("INT", {"default": 4, "min": 2, "max": 1024, "step": 2, "tooltip": "LoRA rank"}), - "save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn"], {"default": "bf16", "tooltip": "the dtype to save the LoRA as"}), + "save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "bf16", "tooltip": "the dtype to save the LoRA as"}), "load_device": (["cpu", "cuda"], {"default": "cuda", "tooltip": "the device to load the model to"}), "store_device": (["cpu", "cuda"], {"default": "cpu", "tooltip": "the device to store the LoRA as"}), "clamp_quantile": ("FLOAT", {"default": 0.99, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "clamp quantile"}), @@ -1427,7 +1494,8 @@ NODE_CLASS_MAPPINGS = { "FluxTrainSaveModel": FluxTrainSaveModel, "ExtractFluxLoRA": ExtractFluxLoRA, "OptimizerConfigProdigy": OptimizerConfigProdigy, - "FluxTrainResume": FluxTrainResume + "FluxTrainResume": FluxTrainResume, + "FluxTrainBlockSelect": FluxTrainBlockSelect } NODE_DISPLAY_NAME_MAPPINGS = { "InitFluxLoRATraining": "Init Flux LoRA Training", @@ -1448,5 +1516,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FluxTrainSaveModel": "Flux Train Save Model", "ExtractFluxLoRA": "Extract Flux LoRA", "OptimizerConfigProdigy": "Optimizer Config Prodigy", - "FluxTrainResume": "Flux Train Resume" + "FluxTrainResume": "Flux Train Resume", + "FluxTrainBlockSelect": "Flux Train Block Select" } diff --git a/train_network.py b/train_network.py index 1950b62..6cb0950 100644 --- a/train_network.py +++ b/train_network.py @@ -158,6 +158,9 @@ class NetworkTrainer: # region SD/SDXL + def post_process_network(self, args, accelerator, network, text_encoders, unet): + pass + def get_noise_scheduler(self, args: argparse.Namespace, device: torch.device) -> Any: noise_scheduler = DDPMScheduler( beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000, clip_sample=False @@ -231,13 +234,20 @@ class NetworkTrainer: def get_sai_model_spec(self, args): return train_util.get_sai_model_spec(None, args, self.is_sdxl, True, False) - - def is_text_encoder_not_needed_for_training(self, args): - return False # use for sample images def update_metadata(self, metadata, args): pass + def is_text_encoder_not_needed_for_training(self, args): + return False # use for sample images + + def prepare_text_encoder_grad_ckpt_workaround(self, index, text_encoder): + # set top parameter requires_grad = True for gradient checkpointing works + text_encoder.text_model.embeddings.requires_grad_(True) + + def prepare_text_encoder_fp8(self, index, text_encoder, te_weight_dtype, weight_dtype): + text_encoder.text_model.embeddings.to(dtype=weight_dtype) + # endregion def init_train(self, args): @@ -318,7 +328,7 @@ class NetworkTrainer: collator = train_util.collator_class(current_epoch, current_step, ds_for_collator) if args.debug_dataset: - train_dataset_group.set_current_strategies() + train_dataset_group.set_current_strategies() # dasaset needs to know the strategies explicitly train_util.debug_dataset(train_dataset_group) return if len(train_dataset_group) == 0: @@ -332,7 +342,7 @@ class NetworkTrainer: train_dataset_group.is_latent_cacheable() ), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません" - self.assert_extra_args(args, train_dataset_group) + self.assert_extra_args(args, train_dataset_group) # may change some args # prepare accelerator logger.info("preparing accelerator") @@ -434,12 +444,15 @@ class NetworkTrainer: ) args.scale_weight_norms = False + self.post_process_network(args, accelerator, network, text_encoders, unet) + + # apply network to unet and text_encoder train_unet = not args.network_train_text_encoder_only train_text_encoder = self.is_train_text_encoder(args) network.apply_to(text_encoder, unet, train_text_encoder, train_unet) if args.network_weights is not None: - # FIXME consider alpha of weights + # FIXME consider alpha of weights: this assumes that the alpha is not changed info = network.load_weights(args.network_weights) accelerator.print(f"load network weights from {args.network_weights}: {info}") @@ -542,11 +555,12 @@ class NetworkTrainer: args.mixed_precision != "no" ), "fp8_base requires mixed precision='fp16' or 'bf16'" accelerator.print("enable fp8 training for U-Net.") - unet_weight_dtype = torch.float8_e4m3fn + unet_weight_dtype = torch.float8_e4m3fn if args.fp8_dtype == "e4m3" else torch.float8_e5m2 + accelerator.print(f"unet_weight_dtype: {unet_weight_dtype}") if not args.fp8_base_unet and not args.network_train_unet_only: accelerator.print("enable fp8 training for Text Encoder.") - te_weight_dtype = weight_dtype if args.fp8_base_unet else torch.float8_e4m3fn + te_weight_dtype = torch.float8_e4m3fn if args.fp8_dtype == "e4m3" else torch.float8_e5m2 # unet.to(accelerator.device) # this makes faster `to(dtype)` below, but consumes 23 GB VRAM # unet.to(dtype=unet_weight_dtype) # without moving to gpu, this takes a lot of time and main memory @@ -555,17 +569,16 @@ class NetworkTrainer: unet.requires_grad_(False) unet.to(dtype=unet_weight_dtype) - for t_enc in text_encoders: + for i, t_enc in enumerate(text_encoders): t_enc.requires_grad_(False) # in case of cpu, dtype is already set to fp32 because cpu does not support fp8/fp16/bf16 if t_enc.device.type != "cpu": t_enc.to(dtype=te_weight_dtype) - if hasattr(t_enc, "text_model") and hasattr(t_enc.text_model, "embeddings"): - # nn.Embedding not support FP8 - t_enc.text_model.embeddings.to(dtype=(weight_dtype if te_weight_dtype != weight_dtype else te_weight_dtype)) - elif hasattr(t_enc, "encoder") and hasattr(t_enc.encoder, "embeddings"): - t_enc.encoder.embeddings.to(dtype=(weight_dtype if te_weight_dtype != weight_dtype else te_weight_dtype)) + + # nn.Embedding not support FP8 + if te_weight_dtype != weight_dtype: + self.prepare_text_encoder_fp8(i, t_enc, te_weight_dtype, weight_dtype) # acceleratorがなんかよろしくやってくれるらしい / accelerator will do something good if args.deepspeed: @@ -606,12 +619,12 @@ class NetworkTrainer: if args.gradient_checkpointing: # according to TI example in Diffusers, train is required unet.train() - for t_enc, frag in zip(text_encoders, self.get_text_encoders_train_flags(args, text_encoders)): + for i, (t_enc, frag) in enumerate(zip(text_encoders, self.get_text_encoders_train_flags(args, text_encoders))): t_enc.train() # set top parameter requires_grad = True for gradient checkpointing works if frag: - t_enc.text_model.embeddings.requires_grad_(True) + self.prepare_text_encoder_grad_ckpt_workaround(i, t_enc) else: unet.eval() @@ -1036,8 +1049,12 @@ class NetworkTrainer: # log device and dtype for each model logger.info(f"unet dtype: {unet_weight_dtype}, device: {unet.device}") - for t_enc in text_encoders: - logger.info(f"text_encoder dtype: {t_enc.dtype}, device: {t_enc.device}") + for i, t_enc in enumerate(text_encoders): + params_itr = t_enc.parameters() + params_itr.__next__() # skip the first parameter + params_itr.__next__() # skip the second parameter. because CLIP first two parameters are embeddings + param_3rd = params_itr.__next__() + logger.info(f"text_encoder [{i}] dtype: {param_3rd.dtype}, device: {t_enc.device}") clean_memory_on_device(accelerator.device) @@ -1058,10 +1075,13 @@ class NetworkTrainer: self.lr_scheduler = lr_scheduler self.save_model = save_model self.remove_model = remove_model + self.comfy_pbar = None progress_bar = tqdm(range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps") + def training_loop(break_at_steps, epoch): steps_done = 0 + #accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}") progress_bar.set_description(f"Epoch {epoch + 1}/{num_train_epochs} - steps") @@ -1108,15 +1128,11 @@ class NetworkTrainer: # print(f"set multiplier: {multipliers}") accelerator.unwrap_model(network).set_multiplier(multipliers) + text_encoder_conds = [] text_encoder_outputs_list = batch.get("text_encoder_outputs_list", None) if text_encoder_outputs_list is not None: text_encoder_conds = text_encoder_outputs_list # List of text encoder outputs - if ( - text_encoder_conds is None - or len(text_encoder_conds) == 0 - or text_encoder_conds[0] is None - or train_text_encoder - ): + if len(text_encoder_conds) == 0 or text_encoder_conds[0] is None or train_text_encoder: with torch.set_grad_enabled(train_text_encoder), accelerator.autocast(): # Get the text embedding for conditioning if args.weighted_captions: @@ -1139,10 +1155,14 @@ class NetworkTrainer: if args.full_fp16: encoded_text_encoder_conds = [c.to(weight_dtype) for c in encoded_text_encoder_conds] - # if encoded_text_encoder_conds is not None, update cached text_encoder_conds - for i in range(len(encoded_text_encoder_conds)): - if encoded_text_encoder_conds[i] is not None: - text_encoder_conds[i] = encoded_text_encoder_conds[i] + # if text_encoder_conds is not cached, use encoded_text_encoder_conds + if len(text_encoder_conds) == 0: + text_encoder_conds = encoded_text_encoder_conds + else: + # if encoded_text_encoder_conds is not None, update cached text_encoder_conds + for i in range(len(encoded_text_encoder_conds)): + if encoded_text_encoder_conds[i] is not None: + text_encoder_conds[i] = encoded_text_encoder_conds[i] # sample noise, call unet, get target noise_pred, target, timesteps, huber_c, weighting = self.get_noise_pred_and_target( @@ -1217,6 +1237,7 @@ class NetworkTrainer: if self.global_step >= break_at_steps: break steps_done += 1 + self.comfy_pbar.update(1) if args.logging_dir is not None: logs = {"loss/epoch": self.loss_recorder.moving_average} @@ -1270,6 +1291,12 @@ def setup_parser() -> argparse.ArgumentParser: parser.add_argument("--unet_lr", type=float, default=None, help="learning rate for U-Net / U-Netの学習率") parser.add_argument("--text_encoder_lr", type=float, default=None, help="learning rate for Text Encoder / Text Encoderの学習率") + parser.add_argument( + "--fp8_base_unet", + action="store_true", + help="use fp8 for U-Net (or DiT), Text Encoder is fp16 or bf16" + " / U-Net(またはDiT)にfp8を使用する。Text Encoderはfp16またはbf16", + ) parser.add_argument( "--network_weights", type=str, default=None, help="pretrained weights for network / 学習するネットワークの初期重み" @@ -1366,10 +1393,9 @@ def setup_parser() -> argparse.ArgumentParser: + " / 初期ステップ数、全エポックを含むステップ数、0で最初のステップ(未指定時と同じ)。initial_epochを上書きする", ) parser.add_argument( - "--fp8_base_unet", + "--cpu_offload_checkpointing", action="store_true", - help="use fp8 for U-Net (or DiT), Text Encoder is fp16 or bf16" - " / U-Net(またはDiT)にfp8を使用する。Text Encoderはfp16またはbf16", + help="[EXPERIMENTAL] enable offloading of tensors to CPU during checkpointing for U-Net or DiT, if supported", ) parser.add_argument( "--cpu_offload_checkpointing",