From e271bbd8f112fea486f4d2313dd07debcc0cfca2 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 18 Aug 2024 21:40:14 +0300 Subject: [PATCH] updates --- examples/flux_lora_train_test_01.json | 2358 +++++++++++++++++++++++++ flux_trai_orig.py | 809 +++++++++ flux_train_comfy.py | 842 +++++++++ flux_train_network_comfy.py | 61 +- library/flux_models.py | 189 +- library/flux_train_utils.py | 270 ++- library/train_util.py | 2 +- nodes.py | 387 +++- 8 files changed, 4741 insertions(+), 177 deletions(-) create mode 100644 examples/flux_lora_train_test_01.json create mode 100644 flux_trai_orig.py create mode 100644 flux_train_comfy.py diff --git a/examples/flux_lora_train_test_01.json b/examples/flux_lora_train_test_01.json new file mode 100644 index 0000000..28fd766 --- /dev/null +++ b/examples/flux_lora_train_test_01.json @@ -0,0 +1,2358 @@ +{ + "last_node_id": 89, + "last_link_id": 136, + "nodes": [ + { + "id": 2, + "type": "FluxTrainModelSelect", + "pos": [ + 250, + 87 + ], + "size": { + "0": 430, + "1": 130 + }, + "flags": {}, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "flux_models", + "type": "TRAIN_FLUX_MODELS", + "links": [ + 64 + ], + "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": 3, + "type": "TrainDatasetConfig", + "pos": [ + 261, + 264 + ], + "size": { + "0": 420, + "1": 410 + }, + "flags": {}, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "dataset", + "type": "TOML_DATASET", + "links": [ + 65 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "TrainDatasetConfig" + }, + "widgets_values": [ + 1024, + 1024, + 1, + "../datasets/akihiko_yoshida_no_caps", + "", + true, + false, + 256, + "1024, 768, 512", + false, + false, + 1 + ] + }, + { + "id": 4, + "type": "FluxTrainLoop", + "pos": [ + 1529, + 338 + ], + "size": { + "0": 393, + "1": 58 + }, + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 66 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 7 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainLoop" + }, + "widgets_values": [ + 500 + ], + "color": "#232", + "bgcolor": "#353" + }, + { + "id": 8, + "type": "FluxTrainValidate", + "pos": [ + 1538.6797892578106, + 487.55795644042973 + ], + "size": { + "0": 468.5999755859375, + "1": 46 + }, + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 7 + }, + { + "name": "validation_settings", + "type": "VALSETTINGS", + "link": 60 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 40, + 133 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "validation_images", + "type": "IMAGE", + "links": [ + 8, + 112 + ], + "slot_index": 1, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainValidate" + } + }, + { + "id": 9, + "type": "PreviewImage", + "pos": [ + 1528.6797892578106, + 587.5579564404297 + ], + "size": { + "0": 949.7310791015625, + "1": 468.02789306640625 + }, + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 8 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 14, + "type": "FluxTrainSave", + "pos": [ + 2044, + 436 + ], + "size": { + "0": 393, + "1": 98 + }, + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 40 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 72 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "lora_path", + "type": "STRING", + "links": [], + "slot_index": 1, + "shape": 3 + }, + { + "name": "steps", + "type": "INT", + "links": [ + 110 + ], + "slot_index": 2, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainSave" + }, + "widgets_values": [ + false + ] + }, + { + "id": 37, + "type": "FluxTrainValidationSettings", + "pos": [ + 773, + 13 + ], + "size": { + "0": 315, + "1": 178 + }, + "flags": {}, + "order": 7, + "mode": 0, + "outputs": [ + { + "name": "validation_settings", + "type": "VALSETTINGS", + "links": [ + 58 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainValidationSettings" + }, + "widgets_values": [ + 20, + 1024, + 1024, + 3.5, + 42, + "fixed" + ] + }, + { + "id": 38, + "type": "SetNode", + "pos": { + "0": 1137, + "1": 43, + "2": 0, + "3": 0, + "4": 0, + "5": 0, + "6": 0, + "7": 0, + "8": 0, + "9": 0 + }, + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "VALSETTINGS", + "type": "VALSETTINGS", + "link": 58 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": null + } + ], + "title": "Set_validation_settings", + "properties": { + "previousName": "validation_settings" + }, + "widgets_values": [ + "validation_settings" + ] + }, + { + "id": 40, + "type": "GetNode", + "pos": { + "0": 1528.677978515625, + "1": 437.5579528808594, + "2": 0, + "3": 0, + "4": 0, + "5": 0, + "6": 0, + "7": 0, + "8": 0, + "9": 0 + }, + "size": { + "0": 277.0899353027344, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "VALSETTINGS", + "type": "VALSETTINGS", + "links": [ + 60 + ], + "slot_index": 0 + } + ], + "title": "Get_validation_settings", + "properties": {}, + "widgets_values": [ + "validation_settings" + ] + }, + { + "id": 42, + "type": "InitFluxLoRATraining", + "pos": [ + 760, + 280 + ], + "size": { + "0": 437.6565246582031, + "1": 802.5902099609375 + }, + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "flux_models", + "type": "TRAIN_FLUX_MODELS", + "link": 64 + }, + { + "name": "dataset_settings", + "type": "TOML_DATASET", + "link": 65 + }, + { + "name": "optimizer_settings", + "type": "ARGS", + "link": 67 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 66 + ], + "shape": 3 + }, + { + "name": "epochs_count", + "type": "INT", + "links": [ + 134 + ], + "shape": 3, + "slot_index": 1 + }, + { + "name": "output_path", + "type": "STRING", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "InitFluxLoRATraining" + }, + "widgets_values": [ + "flux_lora", + "flux_lora_test", + 4, + 0.0004, + 0.0001, + 1500, + true, + 0.0001, + true, + 512, + "disk", + "disk", + false, + "logit_normal", + 0, + 1, + 1.29, + "sigmoid", + 1, + "raw", + 1, + 1, + false, + true, + "fp32", + "bf16", + "sdpa", + "illustration of a kitten | photograph of a turtle" + ] + }, + { + "id": 43, + "type": "OptimizerConfig", + "pos": [ + 362, + 726 + ], + "size": { + "0": 315, + "1": 178 + }, + "flags": {}, + "order": 3, + "mode": 0, + "outputs": [ + { + "name": "optimizer_settings", + "type": "ARGS", + "links": [ + 67 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "OptimizerConfig" + }, + "widgets_values": [ + "adamw8bit", + 1, + "cosine_with_restarts", + 0, + 1, + 1 + ] + }, + { + "id": 44, + "type": "FluxTrainLoop", + "pos": [ + 2636, + 338 + ], + "size": { + "0": 393, + "1": 58 + }, + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 72 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainLoop" + }, + "widgets_values": [ + 500 + ], + "color": "#232", + "bgcolor": "#353" + }, + { + "id": 45, + "type": "FluxTrainValidate", + "pos": [ + 2640, + 500 + ], + "size": { + "0": 468.5999755859375, + "1": 46 + }, + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 133 + }, + { + "name": "validation_settings", + "type": "VALSETTINGS", + "link": 69 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 71 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "validation_images", + "type": "IMAGE", + "links": [ + 70, + 119 + ], + "slot_index": 1, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainValidate" + } + }, + { + "id": 46, + "type": "PreviewImage", + "pos": [ + 2654, + 609 + ], + "size": { + "0": 850.0181274414062, + "1": 452.6767578125 + }, + "flags": {}, + "order": 19, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 70 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 47, + "type": "FluxTrainSave", + "pos": [ + 3137, + 457 + ], + "size": { + "0": 393, + "1": 98 + }, + "flags": {}, + "order": 18, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 71 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 97 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "lora_path", + "type": "STRING", + "links": null, + "shape": 3 + }, + { + "name": "steps", + "type": "INT", + "links": [ + 116 + ], + "slot_index": 2, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainSave" + }, + "widgets_values": [ + false + ] + }, + { + "id": 48, + "type": "GetNode", + "pos": { + "0": 2630, + "1": 450, + "2": 0, + "3": 0, + "4": 0, + "5": 0, + "6": 0, + "7": 0, + "8": 0, + "9": 0 + }, + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "VALSETTINGS", + "type": "VALSETTINGS", + "links": [ + 69 + ], + "slot_index": 0 + } + ], + "title": "Get_validation_settings", + "properties": {}, + "widgets_values": [ + "validation_settings" + ] + }, + { + "id": 59, + "type": "FluxTrainLoop", + "pos": [ + 3697, + 338 + ], + "size": { + "0": 393, + "1": 58 + }, + "flags": {}, + "order": 21, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 97 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 88 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainLoop" + }, + "widgets_values": [ + 500 + ], + "color": "#232", + "bgcolor": "#353" + }, + { + "id": 60, + "type": "FluxTrainValidate", + "pos": [ + 3716.7086387499994, + 510 + ], + "size": { + "0": 468.5999755859375, + "1": 46 + }, + "flags": {}, + "order": 23, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 88 + }, + { + "name": "validation_settings", + "type": "VALSETTINGS", + "link": 89 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 91 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "validation_images", + "type": "IMAGE", + "links": [ + 90, + 122 + ], + "slot_index": 1, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainValidate" + } + }, + { + "id": 61, + "type": "PreviewImage", + "pos": [ + 3707, + 610 + ], + "size": { + "0": 949.7310791015625, + "1": 468.02789306640625 + }, + "flags": {}, + "order": 26, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 90 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 62, + "type": "FluxTrainSave", + "pos": [ + 4227, + 464 + ], + "size": { + "0": 393, + "1": 98 + }, + "flags": {}, + "order": 25, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 91 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 92 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "lora_path", + "type": "STRING", + "links": null, + "shape": 3 + }, + { + "name": "steps", + "type": "INT", + "links": [ + 120 + ], + "slot_index": 2, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainSave" + }, + "widgets_values": [ + false + ] + }, + { + "id": 63, + "type": "GetNode", + "pos": { + "0": 3706.7109375, + "1": 460, + "2": 0, + "3": 0, + "4": 0, + "5": 0, + "6": 0, + "7": 0, + "8": 0, + "9": 0 + }, + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "VALSETTINGS", + "type": "VALSETTINGS", + "links": [ + 89 + ], + "slot_index": 0 + } + ], + "title": "Get_validation_settings", + "properties": {}, + "widgets_values": [ + "validation_settings" + ] + }, + { + "id": 64, + "type": "FluxTrainLoop", + "pos": [ + 4765, + 358 + ], + "size": { + "0": 393, + "1": 58 + }, + "flags": {}, + "order": 27, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 92 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 93 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainLoop" + }, + "widgets_values": [ + 500 + ], + "color": "#232", + "bgcolor": "#353" + }, + { + "id": 65, + "type": "FluxTrainValidate", + "pos": [ + 4775.216642000014, + 518.4568783310547 + ], + "size": { + "0": 468.5999755859375, + "1": 46 + }, + "flags": {}, + "order": 29, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 93 + }, + { + "name": "validation_settings", + "type": "VALSETTINGS", + "link": 94 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 96 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "validation_images", + "type": "IMAGE", + "links": [ + 95, + 126 + ], + "slot_index": 1, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainValidate" + } + }, + { + "id": 66, + "type": "PreviewImage", + "pos": [ + 4785.216642000014, + 628.4568783310547 + ], + "size": { + "0": 850.0181274414062, + "1": 452.6767578125 + }, + "flags": {}, + "order": 32, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 95 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 67, + "type": "FluxTrainSave", + "pos": [ + 5274, + 482 + ], + "size": { + "0": 393, + "1": 98 + }, + "flags": {}, + "order": 31, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 96 + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": [ + 98, + 99 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "lora_path", + "type": "STRING", + "links": [], + "slot_index": 1, + "shape": 3 + }, + { + "name": "steps", + "type": "INT", + "links": [ + 125 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainSave" + }, + "widgets_values": [ + false + ] + }, + { + "id": 68, + "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 + }, + "size": { + "0": 210, + "1": 58 + }, + "flags": { + "collapsed": true + }, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "VALSETTINGS", + "type": "VALSETTINGS", + "links": [ + 94 + ], + "slot_index": 0 + } + ], + "title": "Get_validation_settings", + "properties": {}, + "widgets_values": [ + "validation_settings" + ] + }, + { + "id": 69, + "type": "FluxTrainEnd", + "pos": [ + 5890, + 500 + ], + "size": { + "0": 317.4000244140625, + "1": 78 + }, + "flags": {}, + "order": 33, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 98 + } + ], + "outputs": [ + { + "name": "lora_path", + "type": "STRING", + "links": [ + 103, + 135 + ], + "slot_index": 0, + "shape": 3 + }, + { + "name": "metadata", + "type": "STRING", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "FluxTrainEnd" + }, + "widgets_values": [ + true + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 70, + "type": "VisualizeLoss", + "pos": [ + 5876, + -181 + ], + "size": { + "0": 254.40000915527344, + "1": 26 + }, + "flags": {}, + "order": 34, + "mode": 0, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": 99 + } + ], + "outputs": [ + { + "name": "plot", + "type": "IMAGE", + "links": [ + 100 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "VisualizeLoss" + } + }, + { + "id": 71, + "type": "PreviewImage", + "pos": [ + 5841, + -84 + ], + "size": { + "0": 579.0567016601562, + "1": 541.9993286132812 + }, + "flags": {}, + "order": 38, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 100 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 73, + "type": "Display Any (rgthree)", + "pos": [ + 6270, + 660 + ], + "size": { + "0": 210, + "1": 76 + }, + "flags": {}, + "order": 40, + "mode": 2, + "inputs": [ + { + "name": "source", + "type": "*", + "link": 136, + "dir": 3 + } + ], + "properties": { + "Node name for S&R": "Display Any (rgthree)" + }, + "widgets_values": [ + "" + ] + }, + { + "id": 74, + "type": "Display Any (rgthree)", + "pos": [ + 6275, + 492 + ], + "size": { + "0": 210, + "1": 76.0000228881836 + }, + "flags": {}, + "order": 36, + "mode": 0, + "inputs": [ + { + "name": "source", + "type": "*", + "link": 103, + "dir": 3 + } + ], + "properties": { + "Node name for S&R": "Display Any (rgthree)" + }, + "widgets_values": [ + "" + ] + }, + { + "id": 77, + "type": "ImageConcatMulti", + "pos": [ + 5486, + 1255 + ], + "size": { + "0": 210, + "1": 190 + }, + "flags": {}, + "order": 41, + "mode": 0, + "inputs": [ + { + "name": "image_1", + "type": "IMAGE", + "link": 113 + }, + { + "name": "image_2", + "type": "IMAGE", + "link": 118 + }, + { + "name": "image_3", + "type": "IMAGE", + "link": 123 + }, + { + "name": "image_4", + "type": "IMAGE", + "link": 127 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 128 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": {}, + "widgets_values": [ + 4, + "down", + false, + null + ] + }, + { + "id": 78, + "type": "AddLabel", + "pos": [ + 2055, + 1201 + ], + "size": { + "0": 315, + "1": 274 + }, + "flags": { + "collapsed": true + }, + "order": 20, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 112 + }, + { + "name": "caption", + "type": "STRING", + "link": null, + "widget": { + "name": "caption" + } + }, + { + "name": "text", + "type": "STRING", + "link": 111, + "widget": { + "name": "text" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 113 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "AddLabel" + }, + "widgets_values": [ + 10, + 2, + 48, + 32, + "white", + "black", + "FreeMono.ttf", + "Text", + "up", + "" + ] + }, + { + "id": 79, + "type": "SomethingToString", + "pos": [ + 1833, + 1206 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": { + "collapsed": true + }, + "order": 17, + "mode": 0, + "inputs": [ + { + "name": "input", + "type": "*", + "link": 110 + } + ], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "links": [ + 111 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "SomethingToString" + }, + "widgets_values": [ + "steps ", + "" + ] + }, + { + "id": 80, + "type": "AddLabel", + "pos": [ + 3030, + 1267 + ], + "size": { + "0": 315, + "1": 274 + }, + "flags": { + "collapsed": true + }, + "order": 24, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 119 + }, + { + "name": "caption", + "type": "STRING", + "link": null, + "widget": { + "name": "caption" + } + }, + { + "name": "text", + "type": "STRING", + "link": 117, + "widget": { + "name": "text" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 118 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "AddLabel" + }, + "widgets_values": [ + 10, + 2, + 48, + 32, + "white", + "black", + "FreeMono.ttf", + "Text", + "up", + "" + ] + }, + { + "id": 81, + "type": "SomethingToString", + "pos": [ + 2760, + 1265 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": { + "collapsed": true + }, + "order": 22, + "mode": 0, + "inputs": [ + { + "name": "input", + "type": "*", + "link": 116 + } + ], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "links": [ + 117 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "SomethingToString" + }, + "widgets_values": [ + "steps ", + "" + ] + }, + { + "id": 82, + "type": "SomethingToString", + "pos": [ + 3916, + 1206 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": { + "collapsed": true + }, + "order": 28, + "mode": 0, + "inputs": [ + { + "name": "input", + "type": "*", + "link": 120 + } + ], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "links": [ + 121 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "SomethingToString" + }, + "widgets_values": [ + "steps ", + "" + ] + }, + { + "id": 83, + "type": "AddLabel", + "pos": [ + 4160, + 1201 + ], + "size": { + "0": 315, + "1": 274 + }, + "flags": { + "collapsed": true + }, + "order": 30, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 122 + }, + { + "name": "caption", + "type": "STRING", + "link": null, + "widget": { + "name": "caption" + } + }, + { + "name": "text", + "type": "STRING", + "link": 121, + "widget": { + "name": "text" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 123 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "AddLabel" + }, + "widgets_values": [ + 10, + 2, + 48, + 32, + "white", + "black", + "FreeMono.ttf", + "Text", + "up", + "" + ] + }, + { + "id": 84, + "type": "SomethingToString", + "pos": [ + 4972, + 1195 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": { + "collapsed": true + }, + "order": 35, + "mode": 0, + "inputs": [ + { + "name": "input", + "type": "*", + "link": 125 + } + ], + "outputs": [ + { + "name": "STRING", + "type": "STRING", + "links": [ + 124 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "SomethingToString" + }, + "widgets_values": [ + "steps ", + "" + ] + }, + { + "id": 85, + "type": "AddLabel", + "pos": [ + 5317, + 1251 + ], + "size": { + "0": 315, + "1": 274 + }, + "flags": { + "collapsed": true + }, + "order": 39, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 126 + }, + { + "name": "caption", + "type": "STRING", + "link": null, + "widget": { + "name": "caption" + } + }, + { + "name": "text", + "type": "STRING", + "link": 124, + "widget": { + "name": "text" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 127 + ], + "slot_index": 0, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "AddLabel" + }, + "widgets_values": [ + 10, + 2, + 48, + 32, + "white", + "black", + "FreeMono.ttf", + "Text", + "up", + "" + ] + }, + { + "id": 86, + "type": "PreviewImage", + "pos": [ + 6542, + 442 + ], + "size": { + "0": 532.0540771484375, + "1": 1101.3922119140625 + }, + "flags": {}, + "order": 42, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 128 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 88, + "type": "Display Any (rgthree)", + "pos": [ + 1142, + 124 + ], + "size": { + "0": 210, + "1": 76 + }, + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "source", + "type": "*", + "link": 134, + "dir": 3 + } + ], + "properties": { + "Node name for S&R": "Display Any (rgthree)" + }, + "widgets_values": [ + "" + ] + }, + { + "id": 89, + "type": "UploadToHuggingFace", + "pos": [ + 5900, + 660 + ], + "size": { + "0": 315, + "1": 178 + }, + "flags": {}, + "order": 37, + "mode": 2, + "inputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "link": null + }, + { + "name": "source_path", + "type": "STRING", + "link": 135, + "widget": { + "name": "source_path" + } + } + ], + "outputs": [ + { + "name": "network_trainer", + "type": "NETWORKTRAINER", + "links": null, + "shape": 3 + }, + { + "name": "status", + "type": "STRING", + "links": [ + 136 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "UploadToHuggingFace" + }, + "widgets_values": [ + "", + "", + "", + true, + "" + ] + } + ], + "links": [ + [ + 7, + 4, + 0, + 8, + 0, + "NETWORKTRAINER" + ], + [ + 8, + 8, + 1, + 9, + 0, + "IMAGE" + ], + [ + 40, + 8, + 0, + 14, + 0, + "NETWORKTRAINER" + ], + [ + 58, + 37, + 0, + 38, + 0, + "*" + ], + [ + 60, + 40, + 0, + 8, + 1, + "VALSETTINGS" + ], + [ + 64, + 2, + 0, + 42, + 0, + "TRAIN_FLUX_MODELS" + ], + [ + 65, + 3, + 0, + 42, + 1, + "TOML_DATASET" + ], + [ + 66, + 42, + 0, + 4, + 0, + "NETWORKTRAINER" + ], + [ + 67, + 43, + 0, + 42, + 2, + "ARGS" + ], + [ + 69, + 48, + 0, + 45, + 1, + "VALSETTINGS" + ], + [ + 70, + 45, + 1, + 46, + 0, + "IMAGE" + ], + [ + 71, + 45, + 0, + 47, + 0, + "NETWORKTRAINER" + ], + [ + 72, + 14, + 0, + 44, + 0, + "NETWORKTRAINER" + ], + [ + 88, + 59, + 0, + 60, + 0, + "NETWORKTRAINER" + ], + [ + 89, + 63, + 0, + 60, + 1, + "VALSETTINGS" + ], + [ + 90, + 60, + 1, + 61, + 0, + "IMAGE" + ], + [ + 91, + 60, + 0, + 62, + 0, + "NETWORKTRAINER" + ], + [ + 92, + 62, + 0, + 64, + 0, + "NETWORKTRAINER" + ], + [ + 93, + 64, + 0, + 65, + 0, + "NETWORKTRAINER" + ], + [ + 94, + 68, + 0, + 65, + 1, + "VALSETTINGS" + ], + [ + 95, + 65, + 1, + 66, + 0, + "IMAGE" + ], + [ + 96, + 65, + 0, + 67, + 0, + "NETWORKTRAINER" + ], + [ + 97, + 47, + 0, + 59, + 0, + "NETWORKTRAINER" + ], + [ + 98, + 67, + 0, + 69, + 0, + "NETWORKTRAINER" + ], + [ + 99, + 67, + 0, + 70, + 0, + "NETWORKTRAINER" + ], + [ + 100, + 70, + 0, + 71, + 0, + "IMAGE" + ], + [ + 103, + 69, + 0, + 74, + 0, + "*" + ], + [ + 110, + 14, + 2, + 79, + 0, + "*" + ], + [ + 111, + 79, + 0, + 78, + 2, + "STRING" + ], + [ + 112, + 8, + 1, + 78, + 0, + "IMAGE" + ], + [ + 113, + 78, + 0, + 77, + 0, + "IMAGE" + ], + [ + 116, + 47, + 2, + 81, + 0, + "*" + ], + [ + 117, + 81, + 0, + 80, + 2, + "STRING" + ], + [ + 118, + 80, + 0, + 77, + 1, + "IMAGE" + ], + [ + 119, + 45, + 1, + 80, + 0, + "IMAGE" + ], + [ + 120, + 62, + 2, + 82, + 0, + "*" + ], + [ + 121, + 82, + 0, + 83, + 2, + "STRING" + ], + [ + 122, + 60, + 1, + 83, + 0, + "IMAGE" + ], + [ + 123, + 83, + 0, + 77, + 2, + "IMAGE" + ], + [ + 124, + 84, + 0, + 85, + 2, + "STRING" + ], + [ + 125, + 67, + 2, + 84, + 0, + "*" + ], + [ + 126, + 65, + 1, + 85, + 0, + "IMAGE" + ], + [ + 127, + 85, + 0, + 77, + 3, + "IMAGE" + ], + [ + 128, + 77, + 0, + 86, + 0, + "IMAGE" + ], + [ + 133, + 8, + 0, + 45, + 0, + "NETWORKTRAINER" + ], + [ + 134, + 42, + 1, + 88, + 0, + "*" + ], + [ + 135, + 69, + 0, + 89, + 1, + "STRING" + ], + [ + 136, + 89, + 1, + 73, + 0, + "*" + ] + ], + "groups": [ + { + "title": "Train_01", + "bounding": [ + 1439, + 120, + 1107, + 975 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Settings and init", + "bounding": [ + 193, + -145, + 1174, + 1238 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Train_02", + "bounding": [ + 2602, + 124, + 1046, + 975 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Train_03", + "bounding": [ + 3681, + 128, + 1047, + 986 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Train_04", + "bounding": [ + 4753, + 127, + 996, + 989 + ], + "color": "#3f789e", + "font_size": 24 + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6209213230591554, + "offset": [ + -2130.0418407710777, + 78.17081843628341 + ] + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/flux_trai_orig.py b/flux_trai_orig.py new file mode 100644 index 0000000..62022f0 --- /dev/null +++ b/flux_trai_orig.py @@ -0,0 +1,809 @@ +# training with captions + +# Swap blocks between CPU and GPU: +# This implementation is inspired by and based on the work of 2kpr. +# Many thanks to 2kpr for the original concept and implementation of memory-efficient offloading. +# The original idea has been adapted and extended to fit the current project's needs. + +# Key features: +# - CPU offloading during forward and backward passes +# - Use of fused optimizer and grad_hook for efficient gradient processing +# - Per-block fused optimizer instances + +import argparse +import copy +import math +import os +from multiprocessing import Value +from typing import List +import toml + +from tqdm import tqdm + +import torch +from .library.device_utils import init_ipex, clean_memory_on_device + +init_ipex() + +from accelerate.utils import set_seed +from .library import deepspeed_utils, flux_train_utils, flux_utils, strategy_base, strategy_flux +from .library.sd3_train_utils import load_prompts, FlowMatchEulerDiscreteScheduler + +from .library import train_util as train_util + +from .library.utils import setup_logging, add_logging_arguments + +setup_logging() +import logging + +logger = logging.getLogger(__name__) + +from .library import config_util as config_util + +from .library.config_util import ( + ConfigSanitizer, + BlueprintGenerator, +) +from .library.custom_train_functions import apply_masked_loss, add_custom_train_arguments + + +def train(args): + train_util.verify_training_args(args) + train_util.prepare_dataset_args(args, True) + # sdxl_train_util.verify_sdxl_training_args(args) + deepspeed_utils.prepare_deepspeed_args(args) + setup_logging(args, reset=True) + + # assert ( + # not args.weighted_captions + # ), "weighted_captions is not supported currently / weighted_captionsは現在サポートされていません" + 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.cpu_offload_checkpointing and not args.gradient_checkpointing: + logger.warning( + "cpu_offload_checkpointing is enabled, so gradient_checkpointing is also enabled / cpu_offload_checkpointingが有効になっているため、gradient_checkpointingも有効になります" + ) + args.gradient_checkpointing = True + + cache_latents = args.cache_latents + use_dreambooth_method = args.in_json is None + + if args.seed is not None: + set_seed(args.seed) # 乱数系列を初期化する + + # prepare caching strategy: this must be set before preparing dataset. because dataset may use this strategy for initialization. + if args.cache_latents: + latents_caching_strategy = strategy_flux.FluxLatentsCachingStrategy( + args.cache_latents_to_disk, args.vae_batch_size, args.skip_latents_validity_check + ) + strategy_base.LatentsCachingStrategy.set_strategy(latents_caching_strategy) + + # データセットを準備する + if args.dataset_class is None: + blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, True, args.masked_loss, True)) + if args.dataset_config is not None: + logger.info(f"Load dataset config from {args.dataset_config}") + user_config = config_util.load_user_config(args.dataset_config) + ignored = ["train_data_dir", "in_json"] + if any(getattr(args, attr) is not None for attr in ignored): + logger.warning( + "ignore following options because config file is found: {0} / 設定ファイルが利用されるため以下のオプションは無視されます: {0}".format( + ", ".join(ignored) + ) + ) + else: + if use_dreambooth_method: + logger.info("Using DreamBooth method.") + user_config = { + "datasets": [ + { + "subsets": config_util.generate_dreambooth_subsets_config_by_subdirs( + args.train_data_dir, args.reg_data_dir + ) + } + ] + } + else: + logger.info("Training with captions.") + user_config = { + "datasets": [ + { + "subsets": [ + { + "image_dir": args.train_data_dir, + "metadata_file": args.in_json, + } + ] + } + ] + } + + blueprint = blueprint_generator.generate(user_config, args) + train_dataset_group = config_util.generate_dataset_group_by_blueprint(blueprint.dataset_group) + else: + train_dataset_group = train_util.load_arbitrary_dataset(args) + + current_epoch = Value("i", 0) + current_step = Value("i", 0) + ds_for_collator = train_dataset_group if args.max_data_loader_n_workers == 0 else None + collator = train_util.collator_class(current_epoch, current_step, ds_for_collator) + + train_dataset_group.verify_bucket_reso_steps(16) # TODO これでいいか確認 + + if args.debug_dataset: + if args.cache_text_encoder_outputs: + strategy_base.TextEncoderOutputsCachingStrategy.set_strategy( + strategy_flux.FluxTextEncoderOutputsCachingStrategy( + args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False + ) + ) + train_dataset_group.set_current_strategies() + train_util.debug_dataset(train_dataset_group, True) + return + if len(train_dataset_group) == 0: + logger.error( + "No data found. Please verify the metadata file and train_data_dir option. / 画像がありません。メタデータおよびtrain_data_dirオプションを確認してください。" + ) + return + + if cache_latents: + assert ( + train_dataset_group.is_latent_cacheable() + ), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません" + + 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 / text encoderの出力をキャッシュするときはcaption_dropout_rate, shuffle_caption, token_warmup_step, caption_tag_dropout_rateは使えません" + + # acceleratorを準備する + logger.info("prepare accelerator") + accelerator = train_util.prepare_accelerator(args) + + # mixed precisionに対応した型を用意しておき適宜castする + weight_dtype, save_dtype = train_util.prepare_dtype(args) + + # モデルを読み込む + name = "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev" + + # load VAE for caching latents + ae = None + if cache_latents: + ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu") + ae.to(accelerator.device, dtype=weight_dtype) + ae.requires_grad_(False) + ae.eval() + + train_dataset_group.new_cache_latents(ae, accelerator.is_main_process) + + ae.to("cpu") # if no sampling, vae can be deleted + clean_memory_on_device(accelerator.device) + + accelerator.wait_for_everyone() + + # prepare tokenize strategy + if args.t5xxl_max_token_length is None: + if name == "schnell": + t5xxl_max_token_length = 256 + else: + t5xxl_max_token_length = 512 + else: + t5xxl_max_token_length = args.t5xxl_max_token_length + + flux_tokenize_strategy = strategy_flux.FluxTokenizeStrategy(t5xxl_max_token_length) + strategy_base.TokenizeStrategy.set_strategy(flux_tokenize_strategy) + + # load clip_l, t5xxl for caching text encoder outputs + clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu") + t5xxl = flux_utils.load_t5xxl(args.t5xxl, weight_dtype, "cpu") + clip_l.eval() + t5xxl.eval() + clip_l.requires_grad_(False) + t5xxl.requires_grad_(False) + + text_encoding_strategy = strategy_flux.FluxTextEncodingStrategy(args.apply_t5_attn_mask) + strategy_base.TextEncodingStrategy.set_strategy(text_encoding_strategy) + + # cache text encoder outputs + sample_prompts_te_outputs = None + if args.cache_text_encoder_outputs: + # Text Encodes are eval and no grad here + clip_l.to(accelerator.device) + t5xxl.to(accelerator.device) + + text_encoder_caching_strategy = strategy_flux.FluxTextEncoderOutputsCachingStrategy( + args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False, args.apply_t5_attn_mask + ) + strategy_base.TextEncoderOutputsCachingStrategy.set_strategy(text_encoder_caching_strategy) + + with accelerator.autocast(): + train_dataset_group.new_cache_text_encoder_outputs([clip_l, t5xxl], accelerator.is_main_process) + + # cache sample prompt's embeddings to free text encoder's memory + if args.sample_prompts is not None: + logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}") + + tokenize_strategy: strategy_flux.FluxTokenizeStrategy = strategy_base.TokenizeStrategy.get_strategy() + text_encoding_strategy: strategy_flux.FluxTextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy() + + prompts = load_prompts(args.sample_prompts) + sample_prompts_te_outputs = {} # key: prompt, value: text encoder outputs + with accelerator.autocast(), torch.no_grad(): + for prompt_dict in prompts: + for p in [prompt_dict.get("prompt", ""), prompt_dict.get("negative_prompt", "")]: + if p not in sample_prompts_te_outputs: + logger.info(f"cache Text Encoder outputs for prompt: {p}") + tokens_and_masks = tokenize_strategy.tokenize(p) + sample_prompts_te_outputs[p] = text_encoding_strategy.encode_tokens( + tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, args.apply_t5_attn_mask + ) + + accelerator.wait_for_everyone() + + # now we can delete Text Encoders to free memory + clip_l = None + t5xxl = None + clean_memory_on_device(accelerator.device) + + # load FLUX + # if we load to cpu, flux.to(fp8) takes a long time + flux = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu") + + if args.gradient_checkpointing: + flux.enable_gradient_checkpointing(args.cpu_offload_checkpointing) + + flux.requires_grad_(True) + + if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None: + # Swap blocks between CPU and GPU to reduce memory usage, in forward and backward passes. + # This idea is based on 2kpr's great work. Thank you! + logger.info( + f"enable block swap: double_blocks_to_swap={args.double_blocks_to_swap}, single_blocks_to_swap={args.single_blocks_to_swap}" + ) + flux.enable_block_swap(args.double_blocks_to_swap, args.single_blocks_to_swap) + + if not cache_latents: + # load VAE here if not cached + ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu") + ae.requires_grad_(False) + ae.eval() + ae.to(accelerator.device, dtype=weight_dtype) + + training_models = [] + params_to_optimize = [] + training_models.append(flux) + params_to_optimize.append({"params": list(flux.parameters()), "lr": args.learning_rate}) + + # calculate number of trainable parameters + n_params = 0 + for group in params_to_optimize: + for p in group["params"]: + n_params += p.numel() + + accelerator.print(f"number of trainable parameters: {n_params}") + + # 学習に必要なクラスを準備する + accelerator.print("prepare optimizer, data loader etc.") + + if args.blockwise_fused_optimizers: + # fused backward pass: https://pytorch.org/tutorials/intermediate/optimizer_step_in_backward_tutorial.html + # Instead of creating an optimizer for all parameters as in the tutorial, we create an optimizer for each block of parameters. + # This balances memory usage and management complexity. + + # split params into groups. currently different learning rates are not supported + grouped_params = [] + param_group = {} + for group in params_to_optimize: + named_parameters = list(flux.named_parameters()) + assert len(named_parameters) == len(group["params"]), "number of parameters does not match" + for p, np in zip(group["params"], named_parameters): + # determine target layer and block index for each parameter + block_type = "other" # double, single or other + if np[0].startswith("double_blocks"): + block_idx = int(np[0].split(".")[1]) + block_type = "double" + elif np[0].startswith("single_blocks"): + block_idx = int(np[0].split(".")[1]) + block_type = "single" + else: + block_idx = -1 + + param_group_key = (block_type, block_idx) + if param_group_key not in param_group: + param_group[param_group_key] = [] + param_group[param_group_key].append(p) + + block_types_and_indices = [] + for param_group_key, param_group in param_group.items(): + block_types_and_indices.append(param_group_key) + grouped_params.append({"params": param_group, "lr": args.learning_rate}) + + num_params = 0 + for p in param_group: + num_params += p.numel() + accelerator.print(f"block {param_group_key}: {num_params} parameters") + + # prepare optimizers for each group + optimizers = [] + for group in grouped_params: + _, _, optimizer = train_util.get_optimizer(args, trainable_params=[group]) + optimizers.append(optimizer) + optimizer = optimizers[0] # avoid error in the following code + + logger.info(f"using {len(optimizers)} optimizers for blockwise fused optimizers") + + else: + _, _, optimizer = train_util.get_optimizer(args, trainable_params=params_to_optimize) + + # prepare dataloader + # strategies are set here because they cannot be referenced in another process. Copy them with the dataset + # some strategies can be None + train_dataset_group.set_current_strategies() + + # DataLoaderのプロセス数:0 は persistent_workers が使えないので注意 + n_workers = min(args.max_data_loader_n_workers, os.cpu_count()) # cpu_count or max_data_loader_n_workers + train_dataloader = torch.utils.data.DataLoader( + train_dataset_group, + batch_size=1, + shuffle=True, + collate_fn=collator, + num_workers=n_workers, + persistent_workers=args.persistent_data_loader_workers, + ) + + # 学習ステップ数を計算する + if args.max_train_epochs is not None: + args.max_train_steps = args.max_train_epochs * math.ceil( + len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps + ) + accelerator.print( + f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}" + ) + + # データセット側にも学習ステップを送信 + train_dataset_group.set_max_train_steps(args.max_train_steps) + + # lr schedulerを用意する + if args.blockwise_fused_optimizers: + # prepare lr schedulers for each optimizer + lr_schedulers = [train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes) for optimizer in optimizers] + lr_scheduler = lr_schedulers[0] # avoid error in the following code + else: + lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes) + + # 実験的機能:勾配も含めたfp16/bf16学習を行う モデル全体をfp16/bf16にする + if args.full_fp16: + assert ( + args.mixed_precision == "fp16" + ), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。" + accelerator.print("enable full fp16 training.") + flux.to(weight_dtype) + if clip_l is not None: + clip_l.to(weight_dtype) + t5xxl.to(weight_dtype) # TODO check works with fp16 or not + elif args.full_bf16: + assert ( + args.mixed_precision == "bf16" + ), "full_bf16 requires mixed precision='bf16' / full_bf16を使う場合はmixed_precision='bf16'を指定してください。" + accelerator.print("enable full bf16 training.") + flux.to(weight_dtype) + if clip_l is not None: + clip_l.to(weight_dtype) + t5xxl.to(weight_dtype) + + # if we don't cache text encoder outputs, move them to device + if not args.cache_text_encoder_outputs: + clip_l.to(accelerator.device) + t5xxl.to(accelerator.device) + + clean_memory_on_device(accelerator.device) + + if args.deepspeed: + ds_model = deepspeed_utils.prepare_deepspeed_model(args, mmdit=flux) + # most of ZeRO stage uses optimizer partitioning, so we have to prepare optimizer and ds_model at the same time. # pull/1139#issuecomment-1986790007 + ds_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + ds_model, optimizer, train_dataloader, lr_scheduler + ) + training_models = [ds_model] + + else: + # acceleratorがなんかよろしくやってくれるらしい + flux = accelerator.prepare(flux) + optimizer, train_dataloader, lr_scheduler = accelerator.prepare(optimizer, train_dataloader, lr_scheduler) + + # 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする + if args.full_fp16: + # During deepseed training, accelerate not handles fp16/bf16|mixed precision directly via scaler. Let deepspeed engine do. + # -> But we think it's ok to patch accelerator even if deepspeed is enabled. + train_util.patch_accelerator_for_fp16_training(accelerator) + + # resumeする + train_util.resume_from_local_or_hf_if_specified(accelerator, args) + + if args.fused_backward_pass: + # use fused optimizer for backward pass: other optimizers will be supported in the future + import library.adafactor_fused + + library.adafactor_fused.patch_adafactor_fused(optimizer) + for param_group in optimizer.param_groups: + for parameter in param_group["params"]: + if parameter.requires_grad: + + def __grad_hook(tensor: torch.Tensor, param_group=param_group): + if accelerator.sync_gradients and args.max_grad_norm != 0.0: + accelerator.clip_grad_norm_(tensor, args.max_grad_norm) + optimizer.step_param(tensor, param_group) + tensor.grad = None + + parameter.register_post_accumulate_grad_hook(__grad_hook) + + elif args.blockwise_fused_optimizers: + # prepare for additional optimizers and lr schedulers + for i in range(1, len(optimizers)): + optimizers[i] = accelerator.prepare(optimizers[i]) + lr_schedulers[i] = accelerator.prepare(lr_schedulers[i]) + + # counters are used to determine when to step the optimizer + global optimizer_hooked_count + global num_parameters_per_group + global parameter_optimizer_map + + optimizer_hooked_count = {} + num_parameters_per_group = [0] * len(optimizers) + parameter_optimizer_map = {} + + double_blocks_to_swap = args.double_blocks_to_swap + single_blocks_to_swap = args.single_blocks_to_swap + num_double_blocks = len(flux.double_blocks) + num_single_blocks = len(flux.single_blocks) + + for opt_idx, optimizer in enumerate(optimizers): + for param_group in optimizer.param_groups: + for parameter in param_group["params"]: + if parameter.requires_grad: + block_type, block_idx = block_types_and_indices[opt_idx] + + def create_optimizer_hook(btype, bidx): + def optimizer_hook(parameter: torch.Tensor): + # print(f"optimizer_hook: {btype}, {bidx}") + if accelerator.sync_gradients and args.max_grad_norm != 0.0: + accelerator.clip_grad_norm_(parameter, args.max_grad_norm) + + i = parameter_optimizer_map[parameter] + optimizer_hooked_count[i] += 1 + if optimizer_hooked_count[i] == num_parameters_per_group[i]: + optimizers[i].step() + optimizers[i].zero_grad(set_to_none=True) + + # swap blocks if necessary + if btype == "double" and double_blocks_to_swap: + if bidx >= num_double_blocks - double_blocks_to_swap: + bidx_cuda = double_blocks_to_swap - (num_double_blocks - bidx) + flux.double_blocks[bidx].to("cpu") + flux.double_blocks[bidx_cuda].to(accelerator.device) + # print(f"Move double block {bidx} to cpu and {bidx_cuda} to device") + elif btype == "single" and single_blocks_to_swap: + if bidx >= num_single_blocks - single_blocks_to_swap: + bidx_cuda = single_blocks_to_swap - (num_single_blocks - bidx) + flux.single_blocks[bidx].to("cpu") + flux.single_blocks[bidx_cuda].to(accelerator.device) + # print(f"Move single block {bidx} to cpu and {bidx_cuda} to device") + + return optimizer_hook + + parameter.register_post_accumulate_grad_hook(create_optimizer_hook(block_type, block_idx)) + parameter_optimizer_map[parameter] = opt_idx + num_parameters_per_group[opt_idx] += 1 + + # epoch数を計算する + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + if (args.save_n_epoch_ratio is not None) and (args.save_n_epoch_ratio > 0): + args.save_every_n_epochs = math.floor(num_train_epochs / args.save_n_epoch_ratio) or 1 + + # 学習する + # total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + accelerator.print("running training / 学習開始") + accelerator.print(f" num examples / サンプル数: {train_dataset_group.num_train_images}") + accelerator.print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}") + accelerator.print(f" num epochs / epoch数: {num_train_epochs}") + accelerator.print( + f" batch size per device / バッチサイズ: {', '.join([str(d.batch_size) for d in train_dataset_group.datasets])}" + ) + # accelerator.print( + # f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}" + # ) + accelerator.print(f" gradient accumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}") + accelerator.print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}") + + progress_bar = tqdm(range(args.max_train_steps), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps") + global_step = 0 + + noise_scheduler = FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=args.discrete_flow_shift) + noise_scheduler_copy = copy.deepcopy(noise_scheduler) + + if accelerator.is_main_process: + init_kwargs = {} + if args.wandb_run_name: + init_kwargs["wandb"] = {"name": args.wandb_run_name} + if args.log_tracker_config is not None: + init_kwargs = toml.load(args.log_tracker_config) + accelerator.init_trackers( + "finetuning" if args.log_tracker_name is None else args.log_tracker_name, + config=train_util.get_sanitized_config_or_none(args), + init_kwargs=init_kwargs, + ) + + if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None: + flux.prepare_block_swap_before_forward() + + # For --sample_at_first + #flux_train_utils.sample_images(accelerator, args, 0, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs) + + loss_recorder = train_util.LossRecorder() + epoch = 0 # avoid error when max_train_steps is 0 + for epoch in range(num_train_epochs): + accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}") + current_epoch.value = epoch + 1 + + for m in training_models: + m.train() + + for step, batch in enumerate(train_dataloader): + current_step.value = global_step + + if args.blockwise_fused_optimizers: + optimizer_hooked_count = {i: 0 for i in range(len(optimizers))} # reset counter for each step + + with accelerator.accumulate(*training_models): + if "latents" in batch and batch["latents"] is not None: + latents = batch["latents"].to(accelerator.device, dtype=weight_dtype) + else: + with torch.no_grad(): + # encode images to latents. images are [-1, 1] + latents = ae.encode(batch["images"]) + + # NaNが含まれていれば警告を表示し0に置き換える + if torch.any(torch.isnan(latents)): + accelerator.print("NaN found in latents, replacing with zeros") + latents = torch.nan_to_num(latents, 0, out=latents) + + 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 + else: + # not cached or training, so get from text encoders + tokens_and_masks = batch["input_ids_list"] + with torch.no_grad(): + input_ids = [ids.to(accelerator.device) for ids in batch["input_ids_list"]] + text_encoder_conds = text_encoding_strategy.encode_tokens( + tokenize_strategy, [clip_l, t5xxl], input_ids, args.apply_t5_attn_mask + ) + if args.full_fp16: + text_encoder_conds = [c.to(weight_dtype) for c in text_encoder_conds] + + # TODO support some features for noise implemented in get_noise_noisy_latents_and_timesteps + + # Sample noise that we'll add to the latents + noise = torch.randn_like(latents) + bsz = latents.shape[0] + + # get noisy model input and timesteps + noisy_model_input, timesteps, sigmas = flux_train_utils.get_noisy_model_input_and_timesteps( + args, noise_scheduler, latents, noise, accelerator.device, weight_dtype + ) + + # pack latents and get img_ids + packed_noisy_model_input = flux_utils.pack_latents(noisy_model_input) # b, c, h*2, w*2 -> b, h*w, c*4 + packed_latent_height, packed_latent_width = noisy_model_input.shape[2] // 2, noisy_model_input.shape[3] // 2 + 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) + + # call model + l_pooled, t5_out, txt_ids = text_encoder_conds + with accelerator.autocast(): + # YiYi notes: divide it by 1000 for now because we scale it by 1000 in the transformer model (we should not keep it but I want to keep the inputs same for the model for testing) + model_pred = flux( + img=packed_noisy_model_input, + img_ids=img_ids, + txt=t5_out, + txt_ids=txt_ids, + y=l_pooled, + timesteps=timesteps / 1000, + guidance=guidance_vec, + ) + + # unpack latents + model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width) + + # apply model prediction type + model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas) + + # flow matching loss: this is different from SD3 + target = noise - latents + + # calculate loss + loss = train_util.conditional_loss( + model_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=None + ) + if weighting is not None: + loss = loss * weighting + if args.masked_loss or ("alpha_masks" in batch and batch["alpha_masks"] is not None): + loss = apply_masked_loss(loss, batch) + loss = loss.mean([1, 2, 3]) + + loss_weights = batch["loss_weights"] # 各sampleごとのweight + loss = loss * loss_weights + loss = loss.mean() + + # backward + accelerator.backward(loss) + + if not (args.fused_backward_pass or args.blockwise_fused_optimizers): + if accelerator.sync_gradients and args.max_grad_norm != 0.0: + params_to_clip = [] + for m in training_models: + params_to_clip.extend(m.parameters()) + accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm) + + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad(set_to_none=True) + else: + # optimizer.step() and optimizer.zero_grad() are called in the optimizer hook + lr_scheduler.step() + if args.blockwise_fused_optimizers: + for i in range(1, len(optimizers)): + lr_schedulers[i].step() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + + flux_train_utils.sample_images( + accelerator, args, None, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs + ) + + # 指定ステップごとにモデルを保存 + if args.save_every_n_steps is not None and global_step % args.save_every_n_steps == 0: + accelerator.wait_for_everyone() + if accelerator.is_main_process: + flux_train_utils.save_flux_model_on_epoch_end_or_stepwise( + args, + False, + accelerator, + save_dtype, + epoch, + num_train_epochs, + global_step, + accelerator.unwrap_model(flux), + ) + + current_loss = loss.detach().item() # 平均なのでbatch sizeは関係ないはず + if args.logging_dir is not None: + logs = {"loss": current_loss} + train_util.append_lr_to_logs(logs, lr_scheduler, args.optimizer_type, including_unet=True) + + accelerator.log(logs, step=global_step) + + loss_recorder.add(epoch=epoch, step=step, loss=current_loss) + avr_loss: float = loss_recorder.moving_average + logs = {"avr_loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if args.logging_dir is not None: + logs = {"loss/epoch": loss_recorder.moving_average} + accelerator.log(logs, step=epoch + 1) + + accelerator.wait_for_everyone() + + if args.save_every_n_epochs is not None: + if accelerator.is_main_process: + flux_train_utils.save_flux_model_on_epoch_end_or_stepwise( + args, + True, + accelerator, + save_dtype, + epoch, + num_train_epochs, + global_step, + accelerator.unwrap_model(flux), + ) + + flux_train_utils.sample_images( + accelerator, args, epoch + 1, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs + ) + + is_main_process = accelerator.is_main_process + # if is_main_process: + flux = accelerator.unwrap_model(flux) + + accelerator.end_training() + + if args.save_state or args.save_state_on_train_end: + train_util.save_state_on_train_end(args, accelerator) + + del accelerator # この後メモリを使うのでこれは消す + + if is_main_process: + flux_train_utils.save_flux_model_on_train_end(args, save_dtype, epoch, global_step, flux) + logger.info("model saved.") + + +def setup_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser() + + add_logging_arguments(parser) + train_util.add_sd_models_arguments(parser) # TODO split this + train_util.add_dataset_arguments(parser, True, True, True) + train_util.add_training_arguments(parser, False) + train_util.add_masked_loss_arguments(parser) + deepspeed_utils.add_deepspeed_arguments(parser) + train_util.add_sd_saving_arguments(parser) + train_util.add_optimizer_arguments(parser) + config_util.add_config_arguments(parser) + add_custom_train_arguments(parser) # TODO remove this from here + flux_train_utils.add_flux_train_arguments(parser) + + parser.add_argument( + "--fused_optimizer_groups", + type=int, + default=None, + help="**this option is not working** will be removed in the future / このオプションは動作しません。将来削除されます", + ) + parser.add_argument( + "--blockwise_fused_optimizers", + action="store_true", + help="enable blockwise optimizers for fused backward pass and optimizer step / fused backward passとoptimizer step のためブロック単位のoptimizerを有効にする", + ) + parser.add_argument( + "--skip_latents_validity_check", + action="store_true", + help="skip latents validity check / latentsの正当性チェックをスキップする", + ) + parser.add_argument( + "--double_blocks_to_swap", + type=int, + default=None, + help="[EXPERIMENTAL] " + "Sets the number of 'double_blocks' (~640MB) to swap during the forward and backward passes." + "Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)." + " / 順伝播および逆伝播中にスワップする'変換ブロック'(約640MB)の数を設定します。" + "この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。", + ) + parser.add_argument( + "--single_blocks_to_swap", + type=int, + default=None, + help="[EXPERIMENTAL] " + "Sets the number of 'single_blocks' (~320MB) to swap during the forward and backward passes." + "Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)." + " / 順伝播および逆伝播中にスワップする'変換ブロック'(約320MB)の数を設定します。" + "この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。", + ) + parser.add_argument( + "--cpu_offload_checkpointing", + action="store_true", + help="[EXPERIMENTAL] enable offloading of tensors to CPU during checkpointing / チェックポイント時にテンソルをCPUにオフロードする", + ) + 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) + + train(args) diff --git a/flux_train_comfy.py b/flux_train_comfy.py new file mode 100644 index 0000000..f0b86d0 --- /dev/null +++ b/flux_train_comfy.py @@ -0,0 +1,842 @@ +# training with captions + +# Swap blocks between CPU and GPU: +# This implementation is inspired by and based on the work of 2kpr. +# Many thanks to 2kpr for the original concept and implementation of memory-efficient offloading. +# The original idea has been adapted and extended to fit the current project's needs. + +# Key features: +# - CPU offloading during forward and backward passes +# - Use of fused optimizer and grad_hook for efficient gradient processing +# - Per-block fused optimizer instances + +import argparse +import copy +import math +import os +from multiprocessing import Value +from typing import List +import toml + +from tqdm import tqdm + +import torch +from .library.device_utils import init_ipex, clean_memory_on_device + +init_ipex() + +from accelerate.utils import set_seed +from .library import deepspeed_utils, flux_train_utils, flux_utils, strategy_base, strategy_flux +from .library.sd3_train_utils import load_prompts, FlowMatchEulerDiscreteScheduler + +from .library import train_util as train_util + +from .library.utils import setup_logging, add_logging_arguments + +setup_logging() +import logging + +logger = logging.getLogger(__name__) + +from .library import config_util as config_util + +from .library.config_util import ( + ConfigSanitizer, + BlueprintGenerator, +) +from .library.custom_train_functions import apply_masked_loss, add_custom_train_arguments + + +class FluxTrainer: + def __init__(self): + self.sample_prompts_te_outputs = None + def init_train(self, args): + train_util.verify_training_args(args) + train_util.prepare_dataset_args(args, True) + # sdxl_train_util.verify_sdxl_training_args(args) + deepspeed_utils.prepare_deepspeed_args(args) + setup_logging(args, reset=True) + + # assert ( + # not args.weighted_captions + # ), "weighted_captions is not supported currently / weighted_captionsは現在サポートされていません" + 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.cpu_offload_checkpointing and not args.gradient_checkpointing: + logger.warning( + "cpu_offload_checkpointing is enabled, so gradient_checkpointing is also enabled / cpu_offload_checkpointingが有効になっているため、gradient_checkpointingも有効になります" + ) + args.gradient_checkpointing = True + + cache_latents = args.cache_latents + use_dreambooth_method = args.in_json is None + + if args.seed is not None: + set_seed(args.seed) # 乱数系列を初期化する + + # prepare caching strategy: this must be set before preparing dataset. because dataset may use this strategy for initialization. + if args.cache_latents: + latents_caching_strategy = strategy_flux.FluxLatentsCachingStrategy( + args.cache_latents_to_disk, args.vae_batch_size, args.skip_latents_validity_check + ) + strategy_base.LatentsCachingStrategy.set_strategy(latents_caching_strategy) + + # データセットを準備する + if args.dataset_class is None: + blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, True, args.masked_loss, True)) + if args.dataset_config is not None: + logger.info(f"Load dataset config from {args.dataset_config}") + user_config = config_util.load_user_config(args.dataset_config) + ignored = ["train_data_dir", "in_json"] + if any(getattr(args, attr) is not None for attr in ignored): + logger.warning( + "ignore following options because config file is found: {0} / 設定ファイルが利用されるため以下のオプションは無視されます: {0}".format( + ", ".join(ignored) + ) + ) + else: + if use_dreambooth_method: + logger.info("Using DreamBooth method.") + user_config = { + "datasets": [ + { + "subsets": config_util.generate_dreambooth_subsets_config_by_subdirs( + args.train_data_dir, args.reg_data_dir + ) + } + ] + } + else: + logger.info("Training with captions.") + user_config = { + "datasets": [ + { + "subsets": [ + { + "image_dir": args.train_data_dir, + "metadata_file": args.in_json, + } + ] + } + ] + } + + blueprint = blueprint_generator.generate(user_config, args) + train_dataset_group = config_util.generate_dataset_group_by_blueprint(blueprint.dataset_group) + else: + train_dataset_group = train_util.load_arbitrary_dataset(args) + + current_epoch = Value("i", 0) + current_step = Value("i", 0) + ds_for_collator = train_dataset_group if args.max_data_loader_n_workers == 0 else None + collator = train_util.collator_class(current_epoch, current_step, ds_for_collator) + + train_dataset_group.verify_bucket_reso_steps(16) # TODO これでいいか確認 + + if args.debug_dataset: + if args.cache_text_encoder_outputs: + strategy_base.TextEncoderOutputsCachingStrategy.set_strategy( + strategy_flux.FluxTextEncoderOutputsCachingStrategy( + args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False + ) + ) + train_dataset_group.set_current_strategies() + train_util.debug_dataset(train_dataset_group, True) + return + if len(train_dataset_group) == 0: + logger.error( + "No data found. Please verify the metadata file and train_data_dir option. / 画像がありません。メタデータおよびtrain_data_dirオプションを確認してください。" + ) + return + + if cache_latents: + assert ( + train_dataset_group.is_latent_cacheable() + ), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません" + + 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 / text encoderの出力をキャッシュするときはcaption_dropout_rate, shuffle_caption, token_warmup_step, caption_tag_dropout_rateは使えません" + + # acceleratorを準備する + logger.info("prepare accelerator") + accelerator = train_util.prepare_accelerator(args) + + # mixed precisionに対応した型を用意しておき適宜castする + weight_dtype, save_dtype = train_util.prepare_dtype(args) + + # モデルを読み込む + name = "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev" + + # load VAE for caching latents + ae = None + if cache_latents: + ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu") + ae.to(accelerator.device, dtype=weight_dtype) + ae.requires_grad_(False) + ae.eval() + + train_dataset_group.new_cache_latents(ae, accelerator.is_main_process) + + ae.to("cpu") # if no sampling, vae can be deleted + clean_memory_on_device(accelerator.device) + + accelerator.wait_for_everyone() + + # prepare tokenize strategy + if args.t5xxl_max_token_length is None: + if name == "schnell": + t5xxl_max_token_length = 256 + else: + t5xxl_max_token_length = 512 + else: + t5xxl_max_token_length = args.t5xxl_max_token_length + + flux_tokenize_strategy = strategy_flux.FluxTokenizeStrategy(t5xxl_max_token_length) + strategy_base.TokenizeStrategy.set_strategy(flux_tokenize_strategy) + + # load clip_l, t5xxl for caching text encoder outputs + clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu") + t5xxl = flux_utils.load_t5xxl(args.t5xxl, weight_dtype, "cpu") + clip_l.eval() + t5xxl.eval() + clip_l.requires_grad_(False) + t5xxl.requires_grad_(False) + + text_encoding_strategy = strategy_flux.FluxTextEncodingStrategy(args.apply_t5_attn_mask) + strategy_base.TextEncodingStrategy.set_strategy(text_encoding_strategy) + + # cache text encoder outputs + sample_prompts_te_outputs = None + if args.cache_text_encoder_outputs: + # Text Encodes are eval and no grad here + clip_l.to(accelerator.device, dtype=weight_dtype) + t5xxl.to(accelerator.device, dtype=weight_dtype) + + text_encoder_caching_strategy = strategy_flux.FluxTextEncoderOutputsCachingStrategy( + args.cache_text_encoder_outputs_to_disk, args.text_encoder_batch_size, False, False, args.apply_t5_attn_mask + ) + strategy_base.TextEncoderOutputsCachingStrategy.set_strategy(text_encoder_caching_strategy) + + with accelerator.autocast(): + train_dataset_group.new_cache_text_encoder_outputs([clip_l, t5xxl], accelerator.is_main_process) + + # cache sample prompt's embeddings to free text encoder's memory + if args.sample_prompts is not None: + logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}") + + tokenize_strategy: strategy_flux.FluxTokenizeStrategy = strategy_base.TokenizeStrategy.get_strategy() + text_encoding_strategy: strategy_flux.FluxTextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy() + + prompts = [] + for line in args.sample_prompts: + line = line.strip() + if len(line) > 0 and line[0] != "#": + prompts.append(line) + + # preprocess prompts + for i in range(len(prompts)): + prompt_dict = prompts[i] + if isinstance(prompt_dict, str): + from .library.train_util import line_to_prompt_dict + + prompt_dict = line_to_prompt_dict(prompt_dict) + prompts[i] = prompt_dict + assert isinstance(prompt_dict, dict) + + # Adds an enumerator to the dict based on prompt position. Used later to name image files. Also cleanup of extra data in original prompt dict. + prompt_dict["enum"] = i + prompt_dict.pop("subset", None) + + sample_prompts_te_outputs = {} # key: prompt, value: text encoder outputs + with accelerator.autocast(), torch.no_grad(): + for prompt_dict in prompts: + for p in [prompt_dict.get("prompt", ""), prompt_dict.get("negative_prompt", "")]: + if p not in sample_prompts_te_outputs: + logger.info(f"cache Text Encoder outputs for prompt: {p}") + tokens_and_masks = tokenize_strategy.tokenize(p) + sample_prompts_te_outputs[p] = text_encoding_strategy.encode_tokens( + tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, args.apply_t5_attn_mask + ) + self.sample_prompts_te_outputs = sample_prompts_te_outputs + accelerator.wait_for_everyone() + + # now we can delete Text Encoders to free memory + clip_l = None + t5xxl = None + clean_memory_on_device(accelerator.device) + + # load FLUX + # if we load to cpu, flux.to(fp8) takes a long time + flux = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu") + + if args.gradient_checkpointing: + flux.enable_gradient_checkpointing(args.cpu_offload_checkpointing) + + flux.requires_grad_(True) + + if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None: + # Swap blocks between CPU and GPU to reduce memory usage, in forward and backward passes. + # This idea is based on 2kpr's great work. Thank you! + logger.info( + f"enable block swap: double_blocks_to_swap={args.double_blocks_to_swap}, single_blocks_to_swap={args.single_blocks_to_swap}" + ) + flux.enable_block_swap(args.double_blocks_to_swap, args.single_blocks_to_swap) + + if not cache_latents: + # load VAE here if not cached + ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu") + ae.requires_grad_(False) + ae.eval() + ae.to(accelerator.device, dtype=weight_dtype) + + training_models = [] + params_to_optimize = [] + training_models.append(flux) + params_to_optimize.append({"params": list(flux.parameters()), "lr": args.learning_rate}) + + # calculate number of trainable parameters + n_params = 0 + for group in params_to_optimize: + for p in group["params"]: + n_params += p.numel() + + accelerator.print(f"number of trainable parameters: {n_params}") + + # 学習に必要なクラスを準備する + accelerator.print("prepare optimizer, data loader etc.") + + if args.blockwise_fused_optimizers: + # fused backward pass: https://pytorch.org/tutorials/intermediate/optimizer_step_in_backward_tutorial.html + # Instead of creating an optimizer for all parameters as in the tutorial, we create an optimizer for each block of parameters. + # This balances memory usage and management complexity. + + # split params into groups. currently different learning rates are not supported + grouped_params = [] + param_group = {} + for group in params_to_optimize: + named_parameters = list(flux.named_parameters()) + assert len(named_parameters) == len(group["params"]), "number of parameters does not match" + for p, np in zip(group["params"], named_parameters): + # determine target layer and block index for each parameter + block_type = "other" # double, single or other + if np[0].startswith("double_blocks"): + block_idx = int(np[0].split(".")[1]) + block_type = "double" + elif np[0].startswith("single_blocks"): + block_idx = int(np[0].split(".")[1]) + block_type = "single" + else: + block_idx = -1 + + param_group_key = (block_type, block_idx) + if param_group_key not in param_group: + param_group[param_group_key] = [] + param_group[param_group_key].append(p) + + block_types_and_indices = [] + for param_group_key, param_group in param_group.items(): + block_types_and_indices.append(param_group_key) + grouped_params.append({"params": param_group, "lr": args.learning_rate}) + + num_params = 0 + for p in param_group: + num_params += p.numel() + accelerator.print(f"block {param_group_key}: {num_params} parameters") + + # prepare optimizers for each group + optimizers = [] + for group in grouped_params: + _, _, optimizer = train_util.get_optimizer(args, trainable_params=[group]) + optimizers.append(optimizer) + optimizer = optimizers[0] # avoid error in the following code + + logger.info(f"using {len(optimizers)} optimizers for blockwise fused optimizers") + + else: + _, _, optimizer = train_util.get_optimizer(args, trainable_params=params_to_optimize) + + # prepare dataloader + # strategies are set here because they cannot be referenced in another process. Copy them with the dataset + # some strategies can be None + train_dataset_group.set_current_strategies() + + # DataLoaderのプロセス数:0 は persistent_workers が使えないので注意 + n_workers = min(args.max_data_loader_n_workers, os.cpu_count()) # cpu_count or max_data_loader_n_workers + train_dataloader = torch.utils.data.DataLoader( + train_dataset_group, + batch_size=1, + shuffle=True, + collate_fn=collator, + num_workers=n_workers, + persistent_workers=args.persistent_data_loader_workers, + ) + + # 学習ステップ数を計算する + if args.max_train_epochs is not None: + args.max_train_steps = args.max_train_epochs * math.ceil( + len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps + ) + accelerator.print( + f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}" + ) + + # データセット側にも学習ステップを送信 + train_dataset_group.set_max_train_steps(args.max_train_steps) + + # lr schedulerを用意する + if args.blockwise_fused_optimizers: + # prepare lr schedulers for each optimizer + lr_schedulers = [train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes) for optimizer in optimizers] + lr_scheduler = lr_schedulers[0] # avoid error in the following code + else: + lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes) + + # 実験的機能:勾配も含めたfp16/bf16学習を行う モデル全体をfp16/bf16にする + if args.full_fp16: + assert ( + args.mixed_precision == "fp16" + ), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。" + accelerator.print("enable full fp16 training.") + flux.to(weight_dtype) + if clip_l is not None: + clip_l.to(weight_dtype) + t5xxl.to(weight_dtype) # TODO check works with fp16 or not + elif args.full_bf16: + assert ( + args.mixed_precision == "bf16" + ), "full_bf16 requires mixed precision='bf16' / full_bf16を使う場合はmixed_precision='bf16'を指定してください。" + accelerator.print("enable full bf16 training.") + flux.to(weight_dtype) + if clip_l is not None: + clip_l.to(weight_dtype) + t5xxl.to(weight_dtype) + + # if we don't cache text encoder outputs, move them to device + if not args.cache_text_encoder_outputs: + clip_l.to(accelerator.device) + t5xxl.to(accelerator.device) + + clean_memory_on_device(accelerator.device) + + if args.deepspeed: + ds_model = deepspeed_utils.prepare_deepspeed_model(args, mmdit=flux) + # most of ZeRO stage uses optimizer partitioning, so we have to prepare optimizer and ds_model at the same time. # pull/1139#issuecomment-1986790007 + ds_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + ds_model, optimizer, train_dataloader, lr_scheduler + ) + training_models = [ds_model] + + else: + # acceleratorがなんかよろしくやってくれるらしい + flux = accelerator.prepare(flux) + optimizer, train_dataloader, lr_scheduler = accelerator.prepare(optimizer, train_dataloader, lr_scheduler) + + # 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする + if args.full_fp16: + # During deepseed training, accelerate not handles fp16/bf16|mixed precision directly via scaler. Let deepspeed engine do. + # -> But we think it's ok to patch accelerator even if deepspeed is enabled. + train_util.patch_accelerator_for_fp16_training(accelerator) + + # resumeする + train_util.resume_from_local_or_hf_if_specified(accelerator, args) + + if args.fused_backward_pass: + # use fused optimizer for backward pass: other optimizers will be supported in the future + import library.adafactor_fused + + library.adafactor_fused.patch_adafactor_fused(optimizer) + for param_group in optimizer.param_groups: + for parameter in param_group["params"]: + if parameter.requires_grad: + + def __grad_hook(tensor: torch.Tensor, param_group=param_group): + if accelerator.sync_gradients and args.max_grad_norm != 0.0: + accelerator.clip_grad_norm_(tensor, args.max_grad_norm) + optimizer.step_param(tensor, param_group) + tensor.grad = None + + parameter.register_post_accumulate_grad_hook(__grad_hook) + + elif args.blockwise_fused_optimizers: + # prepare for additional optimizers and lr schedulers + for i in range(1, len(optimizers)): + optimizers[i] = accelerator.prepare(optimizers[i]) + lr_schedulers[i] = accelerator.prepare(lr_schedulers[i]) + + # counters are used to determine when to step the optimizer + global optimizer_hooked_count + global num_parameters_per_group + global parameter_optimizer_map + + optimizer_hooked_count = {} + num_parameters_per_group = [0] * len(optimizers) + parameter_optimizer_map = {} + + double_blocks_to_swap = args.double_blocks_to_swap + single_blocks_to_swap = args.single_blocks_to_swap + num_double_blocks = len(flux.double_blocks) + num_single_blocks = len(flux.single_blocks) + + for opt_idx, optimizer in enumerate(optimizers): + for param_group in optimizer.param_groups: + for parameter in param_group["params"]: + if parameter.requires_grad: + block_type, block_idx = block_types_and_indices[opt_idx] + + def create_optimizer_hook(btype, bidx): + def optimizer_hook(parameter: torch.Tensor): + # print(f"optimizer_hook: {btype}, {bidx}") + if accelerator.sync_gradients and args.max_grad_norm != 0.0: + accelerator.clip_grad_norm_(parameter, args.max_grad_norm) + + i = parameter_optimizer_map[parameter] + optimizer_hooked_count[i] += 1 + if optimizer_hooked_count[i] == num_parameters_per_group[i]: + optimizers[i].step() + optimizers[i].zero_grad(set_to_none=True) + + # swap blocks if necessary + if btype == "double" and double_blocks_to_swap: + if bidx >= num_double_blocks - double_blocks_to_swap: + bidx_cuda = double_blocks_to_swap - (num_double_blocks - bidx) + flux.double_blocks[bidx].to("cpu") + flux.double_blocks[bidx_cuda].to(accelerator.device) + # print(f"Move double block {bidx} to cpu and {bidx_cuda} to device") + elif btype == "single" and single_blocks_to_swap: + if bidx >= num_single_blocks - single_blocks_to_swap: + bidx_cuda = single_blocks_to_swap - (num_single_blocks - bidx) + flux.single_blocks[bidx].to("cpu") + flux.single_blocks[bidx_cuda].to(accelerator.device) + # print(f"Move single block {bidx} to cpu and {bidx_cuda} to device") + + return optimizer_hook + + parameter.register_post_accumulate_grad_hook(create_optimizer_hook(block_type, block_idx)) + parameter_optimizer_map[parameter] = opt_idx + num_parameters_per_group[opt_idx] += 1 + + # epoch数を計算する + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + if (args.save_n_epoch_ratio is not None) and (args.save_n_epoch_ratio > 0): + args.save_every_n_epochs = math.floor(num_train_epochs / args.save_n_epoch_ratio) or 1 + + # 学習する + # total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + accelerator.print("running training / 学習開始") + accelerator.print(f" num examples / サンプル数: {train_dataset_group.num_train_images}") + accelerator.print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}") + accelerator.print(f" num epochs / epoch数: {num_train_epochs}") + accelerator.print( + f" batch size per device / バッチサイズ: {', '.join([str(d.batch_size) for d in train_dataset_group.datasets])}" + ) + # accelerator.print( + # f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}" + # ) + accelerator.print(f" gradient accumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}") + accelerator.print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}") + + progress_bar = tqdm(range(args.max_train_steps), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps") + self.global_step = 0 + + noise_scheduler = FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=args.discrete_flow_shift) + noise_scheduler_copy = copy.deepcopy(noise_scheduler) + + if accelerator.is_main_process: + init_kwargs = {} + if args.wandb_run_name: + init_kwargs["wandb"] = {"name": args.wandb_run_name} + if args.log_tracker_config is not None: + init_kwargs = toml.load(args.log_tracker_config) + accelerator.init_trackers( + "finetuning" if args.log_tracker_name is None else args.log_tracker_name, + config=train_util.get_sanitized_config_or_none(args), + init_kwargs=init_kwargs, + ) + + if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None: + flux.prepare_block_swap_before_forward() + + # For --sample_at_first + #flux_train_utils.sample_images(accelerator, args, 0, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs) + + loss_recorder = train_util.LossRecorder() + epoch = 0 # avoid error when max_train_steps is 0 + + self.tokens_and_masks = tokens_and_masks + self.num_train_epochs = num_train_epochs + self.current_epoch = current_epoch + self.args = args + + def training_loop(break_at_steps, epoch): + global optimizer_hooked_count + steps_done = 0 + accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}") + current_epoch.value = epoch + 1 + + for m in training_models: + m.train() + + for step, batch in enumerate(train_dataloader): + current_step.value = self.global_step + + if args.blockwise_fused_optimizers: + optimizer_hooked_count = {i: 0 for i in range(len(optimizers))} # reset counter for each step + + with accelerator.accumulate(*training_models): + if "latents" in batch and batch["latents"] is not None: + latents = batch["latents"].to(accelerator.device, dtype=weight_dtype) + else: + with torch.no_grad(): + # encode images to latents. images are [-1, 1] + latents = ae.encode(batch["images"]) + + # NaNが含まれていれば警告を表示し0に置き換える + if torch.any(torch.isnan(latents)): + accelerator.print("NaN found in latents, replacing with zeros") + latents = torch.nan_to_num(latents, 0, out=latents) + + 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 + else: + # not cached or training, so get from text encoders + self.tokens_and_masks = batch["input_ids_list"] + with torch.no_grad(): + input_ids = [ids.to(accelerator.device) for ids in batch["input_ids_list"]] + text_encoder_conds = text_encoding_strategy.encode_tokens( + tokenize_strategy, [clip_l, t5xxl], input_ids, args.apply_t5_attn_mask + ) + if args.full_fp16: + text_encoder_conds = [c.to(weight_dtype) for c in text_encoder_conds] + + # TODO support some features for noise implemented in get_noise_noisy_latents_and_timesteps + + # Sample noise that we'll add to the latents + noise = torch.randn_like(latents) + bsz = latents.shape[0] + + # get noisy model input and timesteps + noisy_model_input, timesteps, sigmas = flux_train_utils.get_noisy_model_input_and_timesteps( + args, noise_scheduler, latents, noise, accelerator.device, weight_dtype + ) + + # pack latents and get img_ids + packed_noisy_model_input = flux_utils.pack_latents(noisy_model_input) # b, c, h*2, w*2 -> b, h*w, c*4 + packed_latent_height, packed_latent_width = noisy_model_input.shape[2] // 2, noisy_model_input.shape[3] // 2 + 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) + + # call model + l_pooled, t5_out, txt_ids = text_encoder_conds + with accelerator.autocast(): + # YiYi notes: divide it by 1000 for now because we scale it by 1000 in the transformer model (we should not keep it but I want to keep the inputs same for the model for testing) + model_pred = flux( + img=packed_noisy_model_input, + img_ids=img_ids, + txt=t5_out, + txt_ids=txt_ids, + y=l_pooled, + timesteps=timesteps / 1000, + guidance=guidance_vec, + ) + + # unpack latents + model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width) + + # apply model prediction type + model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas) + + # flow matching loss: this is different from SD3 + target = noise - latents + + # calculate loss + loss = train_util.conditional_loss( + model_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=None + ) + if weighting is not None: + loss = loss * weighting + if args.masked_loss or ("alpha_masks" in batch and batch["alpha_masks"] is not None): + loss = apply_masked_loss(loss, batch) + loss = loss.mean([1, 2, 3]) + + loss_weights = batch["loss_weights"] # 各sampleごとのweight + loss = loss * loss_weights + loss = loss.mean() + + # backward + accelerator.backward(loss) + + if not (args.fused_backward_pass or args.blockwise_fused_optimizers): + if accelerator.sync_gradients and args.max_grad_norm != 0.0: + params_to_clip = [] + for m in training_models: + params_to_clip.extend(m.parameters()) + accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm) + + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad(set_to_none=True) + else: + # optimizer.step() and optimizer.zero_grad() are called in the optimizer hook + lr_scheduler.step() + if args.blockwise_fused_optimizers: + for i in range(1, len(optimizers)): + lr_schedulers[i].step() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + self.global_step += 1 + + # flux_train_utils.sample_images( + # accelerator, args, None, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs + # ) + + # # 指定ステップごとにモデルを保存 + # if args.save_every_n_steps is not None and global_step % args.save_every_n_steps == 0: + # accelerator.wait_for_everyone() + # if accelerator.is_main_process: + # flux_train_utils.save_flux_model_on_epoch_end_or_stepwise( + # args, + # False, + # accelerator, + # save_dtype, + # epoch, + # num_train_epochs, + # global_step, + # accelerator.unwrap_model(flux), + # ) + + current_loss = loss.detach().item() # 平均なのでbatch sizeは関係ないはず + if args.logging_dir is not None: + logs = {"loss": current_loss} + train_util.append_lr_to_logs(logs, lr_scheduler, args.optimizer_type, including_unet=True) + + accelerator.log(logs, step=self.global_step) + + loss_recorder.add(epoch=epoch, step=step, loss=current_loss) + avr_loss: float = loss_recorder.moving_average + logs = {"avr_loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if self.global_step >= break_at_steps: + break + steps_done += 1 + + if args.logging_dir is not None: + logs = {"loss/epoch": loss_recorder.moving_average} + accelerator.log(logs, step=epoch + 1) + return steps_done + + return training_loop + #accelerator.wait_for_everyone() + + # if args.save_every_n_epochs is not None: + # if accelerator.is_main_process: + # flux_train_utils.save_flux_model_on_epoch_end_or_stepwise( + # args, + # True, + # accelerator, + # save_dtype, + # epoch, + # num_train_epochs, + # global_step, + # accelerator.unwrap_model(flux), + # ) + + # flux_train_utils.sample_images( + # accelerator, args, epoch + 1, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs + # ) + + # is_main_process = accelerator.is_main_process + # # if is_main_process: + # flux = accelerator.unwrap_model(flux) + + # accelerator.end_training() + + # if args.save_state or args.save_state_on_train_end: + # train_util.save_state_on_train_end(args, accelerator) + + # del accelerator # この後メモリを使うのでこれは消す + + # if is_main_process: + # flux_train_utils.save_flux_model_on_train_end(args, save_dtype, epoch, global_step, flux) + # logger.info("model saved.") + + +def setup_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser() + + add_logging_arguments(parser) + train_util.add_sd_models_arguments(parser) # TODO split this + train_util.add_dataset_arguments(parser, True, True, True) + train_util.add_training_arguments(parser, False) + train_util.add_masked_loss_arguments(parser) + deepspeed_utils.add_deepspeed_arguments(parser) + train_util.add_sd_saving_arguments(parser) + train_util.add_optimizer_arguments(parser) + config_util.add_config_arguments(parser) + add_custom_train_arguments(parser) # TODO remove this from here + flux_train_utils.add_flux_train_arguments(parser) + + parser.add_argument( + "--fused_optimizer_groups", + type=int, + default=None, + help="**this option is not working** will be removed in the future / このオプションは動作しません。将来削除されます", + ) + parser.add_argument( + "--blockwise_fused_optimizers", + action="store_true", + help="enable blockwise optimizers for fused backward pass and optimizer step / fused backward passとoptimizer step のためブロック単位のoptimizerを有効にする", + ) + parser.add_argument( + "--skip_latents_validity_check", + action="store_true", + help="skip latents validity check / latentsの正当性チェックをスキップする", + ) + parser.add_argument( + "--double_blocks_to_swap", + type=int, + default=None, + help="[EXPERIMENTAL] " + "Sets the number of 'double_blocks' (~640MB) to swap during the forward and backward passes." + "Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)." + " / 順伝播および逆伝播中にスワップする'変換ブロック'(約640MB)の数を設定します。" + "この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。", + ) + parser.add_argument( + "--single_blocks_to_swap", + type=int, + default=None, + help="[EXPERIMENTAL] " + "Sets the number of 'single_blocks' (~320MB) to swap during the forward and backward passes." + "Increasing this number lowers the overall VRAM used during training at the expense of training speed (s/it)." + " / 順伝播および逆伝播中にスワップする'変換ブロック'(約320MB)の数を設定します。" + "この数を増やすと、トレーニング中のVRAM使用量が減りますが、トレーニング速度(s/it)も低下します。", + ) + parser.add_argument( + "--cpu_offload_checkpointing", + action="store_true", + help="[EXPERIMENTAL] enable offloading of tensors to CPU during checkpointing / チェックポイント時にテンソルをCPUにオフロードする", + ) + 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) + +# train(args) diff --git a/flux_train_network_comfy.py b/flux_train_network_comfy.py index 9460883..20f9634 100644 --- a/flux_train_network_comfy.py +++ b/flux_train_network_comfy.py @@ -168,7 +168,7 @@ class FluxNetworkTrainer(NetworkTrainer): for i in range(len(prompts)): prompt_dict = prompts[i] if isinstance(prompt_dict, str): - from library.train_util import line_to_prompt_dict + from .library.train_util import line_to_prompt_dict prompt_dict = line_to_prompt_dict(prompt_dict) prompts[i] = prompt_dict @@ -205,13 +205,8 @@ class FluxNetworkTrainer(NetworkTrainer): text_encoders[0].to(accelerator.device, dtype=weight_dtype) text_encoders[1].to(accelerator.device, dtype=weight_dtype) - def sample_images(self, accelerator, args, epoch, global_step, ae, text_encoder, flux, validation_settings): - if not args.split_mode: - image_tensors = flux_train_utils.sample_images( - accelerator, args, epoch, global_step, flux, ae, text_encoder, self.sample_prompts_te_outputs, validation_settings - ) - return image_tensors - + def sample_images_split_mode(self, accelerator, args, epoch, global_step, flux, ae, text_encoder, sample_prompts_te_outputs, validation_settings): + class FluxUpperLowerWrapper(torch.nn.Module): def __init__(self, flux_upper: flux_models.FluxUpper, flux_lower: flux_models.FluxLower, device: torch.device): super().__init__() @@ -232,7 +227,7 @@ class FluxNetworkTrainer(NetworkTrainer): wrapper = FluxUpperLowerWrapper(self.flux_upper, flux, accelerator.device) clean_memory_on_device(accelerator.device) flux_train_utils.sample_images( - accelerator, args, epoch, global_step, flux, ae, text_encoder, self.sample_prompts_te_outputs, validation_settings + accelerator, args, epoch, global_step, wrapper, ae, text_encoder, sample_prompts_te_outputs, validation_settings ) clean_memory_on_device(accelerator.device) @@ -316,32 +311,10 @@ class FluxNetworkTrainer(NetworkTrainer): noise = torch.randn_like(latents) bsz = latents.shape[0] - if args.timestep_sampling == "uniform" or args.timestep_sampling == "sigmoid": - # Simple random t-based noise sampling - if args.timestep_sampling == "sigmoid": - # https://github.com/XLabs-AI/x-flux/tree/main - t = torch.sigmoid(args.sigmoid_scale * torch.randn((bsz,), device=accelerator.device)) - else: - t = torch.rand((bsz,), device=accelerator.device) - timesteps = t * 1000.0 - t = t.view(-1, 1, 1, 1) - noisy_model_input = (1 - t) * latents + t * noise - else: - # Sample a random timestep for each image - # for weighting schemes where we sample timesteps non-uniformly - u = compute_density_for_timestep_sampling( - weighting_scheme=args.weighting_scheme, - batch_size=bsz, - logit_mean=args.logit_mean, - logit_std=args.logit_std, - mode_scale=args.mode_scale, - ) - indices = (u * self.noise_scheduler_copy.config.num_train_timesteps).long() - timesteps = self.noise_scheduler_copy.timesteps[indices].to(device=accelerator.device) - - # Add noise according to flow matching. - sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=weight_dtype) - noisy_model_input = sigmas * noise + (1.0 - sigmas) * latents + # get noisy model input and timesteps + noisy_model_input, timesteps, sigmas = flux_train_utils.get_noisy_model_input_and_timesteps( + args, noise_scheduler, latents, noise, accelerator.device, weight_dtype + ) # pack latents and get img_ids packed_noisy_model_input = flux_utils.pack_latents(noisy_model_input) # b, c, h*2, w*2 -> b, h*w, c*4 @@ -414,20 +387,8 @@ class FluxNetworkTrainer(NetworkTrainer): # unpack latents model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width) - if args.model_prediction_type == "raw": - # use model_pred as is - weighting = None - elif args.model_prediction_type == "additive": - # add the model_pred to the noisy_model_input - model_pred = model_pred + noisy_model_input - weighting = None - elif args.model_prediction_type == "sigma_scaled": - # apply sigma scaling - model_pred = model_pred * (-sigmas) + noisy_model_input - - # these weighting schemes use a uniform timestep sampling - # and instead post-weight the loss - weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + # apply model prediction type + model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas) # flow matching loss: this is different from SD3 target = noise - latents @@ -450,4 +411,4 @@ class FluxNetworkTrainer(NetworkTrainer): metadata["ss_timestep_sampling"] = args.timestep_sampling 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 \ No newline at end of file + metadata["ss_discrete_flow_shift"] = args.discrete_flow_shift diff --git a/library/flux_models.py b/library/flux_models.py index ef90172..1853bca 100644 --- a/library/flux_models.py +++ b/library/flux_models.py @@ -4,8 +4,12 @@ from dataclasses import dataclass import math - +from typing import Optional import torch +from ..library.device_utils import init_ipex, clean_memory_on_device + +init_ipex() + from einops import rearrange from torch import Tensor, nn from torch.utils.checkpoint import checkpoint @@ -466,6 +470,33 @@ def apply_rope(xq: Tensor, xk: Tensor, freqs_cis: Tensor) -> tuple[Tensor, Tenso # region layers + + +# for cpu_offload_checkpointing + + +def to_cuda(x): + if isinstance(x, torch.Tensor): + return x.cuda() + elif isinstance(x, (list, tuple)): + return [to_cuda(elem) for elem in x] + elif isinstance(x, dict): + return {k: to_cuda(v) for k, v in x.items()} + else: + return x + + +def to_cpu(x): + if isinstance(x, torch.Tensor): + return x.cpu() + elif isinstance(x, (list, tuple)): + return [to_cpu(elem) for elem in x] + elif isinstance(x, dict): + return {k: to_cpu(v) for k, v in x.items()} + else: + return x + + class EmbedND(nn.Module): def __init__(self, dim: int, theta: int, axes_dim: list[int]): super().__init__() @@ -648,16 +679,15 @@ class DoubleStreamBlock(nn.Module): ) self.gradient_checkpointing = False + self.cpu_offload_checkpointing = False - def enable_gradient_checkpointing(self): + def enable_gradient_checkpointing(self, cpu_offload: bool = False): self.gradient_checkpointing = True - # self.img_attn.enable_gradient_checkpointing() - # self.txt_attn.enable_gradient_checkpointing() + self.cpu_offload_checkpointing = cpu_offload def disable_gradient_checkpointing(self): self.gradient_checkpointing = False - # self.img_attn.disable_gradient_checkpointing() - # self.txt_attn.disable_gradient_checkpointing() + self.cpu_offload_checkpointing = False def _forward(self, img: Tensor, txt: Tensor, vec: Tensor, pe: Tensor) -> tuple[Tensor, Tensor]: img_mod1, img_mod2 = self.img_mod(vec) @@ -694,11 +724,24 @@ class DoubleStreamBlock(nn.Module): txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift) return img, txt - def forward(self, *args, **kwargs): + def forward(self, img: Tensor, txt: Tensor, vec: Tensor, pe: Tensor) -> tuple[Tensor, Tensor]: if self.training and self.gradient_checkpointing: - return checkpoint(self._forward, *args, use_reentrant=False, **kwargs) + if not self.cpu_offload_checkpointing: + return checkpoint(self._forward, img, txt, vec, pe, use_reentrant=False) + # cpu offload checkpointing + + def create_custom_forward(func): + def custom_forward(*inputs): + cuda_inputs = to_cuda(inputs) + outputs = func(*cuda_inputs) + return to_cpu(outputs) + + return custom_forward + + return torch.utils.checkpoint.checkpoint(create_custom_forward(self._forward), img, txt, vec, pe) + else: - return self._forward(*args, **kwargs) + return self._forward(img, txt, vec, pe) # def forward(self, img: Tensor, txt: Tensor, vec: Tensor, pe: Tensor): # if self.training and self.gradient_checkpointing: @@ -747,12 +790,15 @@ class SingleStreamBlock(nn.Module): self.modulation = Modulation(hidden_size, double=False) self.gradient_checkpointing = False + self.cpu_offload_checkpointing = False - def enable_gradient_checkpointing(self): + def enable_gradient_checkpointing(self, cpu_offload: bool = False): self.gradient_checkpointing = True + self.cpu_offload_checkpointing = cpu_offload def disable_gradient_checkpointing(self): self.gradient_checkpointing = False + self.cpu_offload_checkpointing = False def _forward(self, x: Tensor, vec: Tensor, pe: Tensor) -> Tensor: mod, _ = self.modulation(vec) @@ -768,11 +814,24 @@ class SingleStreamBlock(nn.Module): output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2)) return x + mod.gate * output - def forward(self, *args, **kwargs): + def forward(self, x: Tensor, vec: Tensor, pe: Tensor) -> Tensor: if self.training and self.gradient_checkpointing: - return checkpoint(self._forward, *args, use_reentrant=False, **kwargs) + if not self.cpu_offload_checkpointing: + return checkpoint(self._forward, x, vec, pe, use_reentrant=False) + + # cpu offload checkpointing + + def create_custom_forward(func): + def custom_forward(*inputs): + cuda_inputs = to_cuda(inputs) + outputs = func(*cuda_inputs) + return to_cpu(outputs) + + return custom_forward + + return torch.utils.checkpoint.checkpoint(create_custom_forward(self._forward), x, vec, pe) else: - return self._forward(*args, **kwargs) + return self._forward(x, vec, pe) # def forward(self, x: Tensor, vec: Tensor, pe: Tensor): # if self.training and self.gradient_checkpointing: @@ -849,6 +908,9 @@ class Flux(nn.Module): self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels) self.gradient_checkpointing = False + self.cpu_offload_checkpointing = False + self.double_blocks_to_swap = None + self.single_blocks_to_swap = None @property def device(self): @@ -858,8 +920,9 @@ class Flux(nn.Module): def dtype(self): return next(self.parameters()).dtype - def enable_gradient_checkpointing(self): + def enable_gradient_checkpointing(self, cpu_offload: bool = False): self.gradient_checkpointing = True + self.cpu_offload_checkpointing = cpu_offload self.time_in.enable_gradient_checkpointing() self.vector_in.enable_gradient_checkpointing() @@ -867,23 +930,42 @@ class Flux(nn.Module): self.guidance_in.enable_gradient_checkpointing() for block in self.double_blocks + self.single_blocks: - block.enable_gradient_checkpointing() + block.enable_gradient_checkpointing(cpu_offload=cpu_offload) - print("FLUX: Gradient checkpointing enabled.") + print(f"FLUX: Gradient checkpointing enabled. CPU offload: {cpu_offload}") def disable_gradient_checkpointing(self): self.gradient_checkpointing = False + self.cpu_offload_checkpointing = False self.time_in.disable_gradient_checkpointing() self.vector_in.disable_gradient_checkpointing() if self.guidance_in.__class__ != nn.Identity: - self.guidance_in.enable_gradient_checkpointing() + self.guidance_in.disable_gradient_checkpointing() for block in self.double_blocks + self.single_blocks: block.disable_gradient_checkpointing() print("FLUX: Gradient checkpointing disabled.") + def enable_block_swap(self, double_blocks: Optional[int], single_blocks: Optional[int]): + self.double_blocks_to_swap = double_blocks + self.single_blocks_to_swap = single_blocks + + def prepare_block_swap_before_forward(self): + # move last n blocks to cpu: they are on cuda + if self.double_blocks_to_swap: + for i in range(len(self.double_blocks) - self.double_blocks_to_swap): + self.double_blocks[i].to(self.device) + for i in range(len(self.double_blocks) - self.double_blocks_to_swap, len(self.double_blocks)): + self.double_blocks[i].to("cpu") # , non_blocking=True) + if self.single_blocks_to_swap: + for i in range(len(self.single_blocks) - self.single_blocks_to_swap): + self.single_blocks[i].to(self.device) + for i in range(len(self.single_blocks) - self.single_blocks_to_swap, len(self.single_blocks)): + self.single_blocks[i].to("cpu") # , non_blocking=True) + clean_memory_on_device(self.device) + def forward( self, img: Tensor, @@ -910,14 +992,75 @@ class Flux(nn.Module): ids = torch.cat((txt_ids, img_ids), dim=1) pe = self.pe_embedder(ids) - for block in self.double_blocks: - img, txt = block(img=img, txt=txt, vec=vec, pe=pe) + if not self.double_blocks_to_swap: + for block in self.double_blocks: + img, txt = block(img=img, txt=txt, vec=vec, pe=pe) + else: + # make sure first n blocks are on cuda, and last n blocks are on cpu at beginning + for block_idx in range(self.double_blocks_to_swap): + block = self.double_blocks[len(self.double_blocks) - self.double_blocks_to_swap + block_idx] + if block.parameters().__next__().device.type != "cpu": + block.to("cpu") # , non_blocking=True) + # print(f"Moved double block {len(self.double_blocks) - self.double_blocks_to_swap + block_idx} to cpu.") + + block = self.double_blocks[block_idx] + if block.parameters().__next__().device.type == "cpu": + block.to(self.device) + # print(f"Moved double block {block_idx} to cuda.") + + to_cpu_block_index = 0 + for block_idx, block in enumerate(self.double_blocks): + # move last n blocks to cuda: they are on cpu, and move first n blocks to cpu: they are on cuda + moving = block_idx >= len(self.double_blocks) - self.double_blocks_to_swap + if moving: + block.to(self.device) # move to cuda + # print(f"Moved double block {block_idx} to cuda.") + + img, txt = block(img=img, txt=txt, vec=vec, pe=pe) + + if moving: + self.double_blocks[to_cpu_block_index].to("cpu") # , non_blocking=True) + # print(f"Moved double block {to_cpu_block_index} to cpu.") + to_cpu_block_index += 1 img = torch.cat((txt, img), 1) - for block in self.single_blocks: - img = block(img, vec=vec, pe=pe) + + if not self.single_blocks_to_swap: + for block in self.single_blocks: + img = block(img, vec=vec, pe=pe) + else: + # make sure first n blocks are on cuda, and last n blocks are on cpu at beginning + for block_idx in range(self.single_blocks_to_swap): + block = self.single_blocks[len(self.single_blocks) - self.single_blocks_to_swap + block_idx] + if block.parameters().__next__().device.type != "cpu": + block.to("cpu") # , non_blocking=True) + # print(f"Moved single block {len(self.single_blocks) - self.single_blocks_to_swap + block_idx} to cpu.") + + block = self.single_blocks[block_idx] + if block.parameters().__next__().device.type == "cpu": + block.to(self.device) + # print(f"Moved single block {block_idx} to cuda.") + + to_cpu_block_index = 0 + for block_idx, block in enumerate(self.single_blocks): + # move last n blocks to cuda: they are on cpu, and move first n blocks to cpu: they are on cuda + moving = block_idx >= len(self.single_blocks) - self.single_blocks_to_swap + if moving: + block.to(self.device) # move to cuda + # print(f"Moved single block {block_idx} to cuda.") + + img = block(img, vec=vec, pe=pe) + + if moving: + self.single_blocks[to_cpu_block_index].to("cpu") # , non_blocking=True) + # print(f"Moved single block {to_cpu_block_index} to cpu.") + img = img[:, txt.shape[1] :, ...] + if self.training and self.cpu_offload_checkpointing: + img = img.to(self.device) + vec = vec.to(self.device) + img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels) return img @@ -988,7 +1131,7 @@ class FluxUpper(nn.Module): self.time_in.disable_gradient_checkpointing() self.vector_in.disable_gradient_checkpointing() if self.guidance_in.__class__ != nn.Identity: - self.guidance_in.enable_gradient_checkpointing() + self.guidance_in.disable_gradient_checkpointing() for block in self.double_blocks: block.disable_gradient_checkpointing() @@ -1086,4 +1229,4 @@ class FluxLower(nn.Module): img = img[:, txt.shape[1] :, ...] img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels) - return img + return img \ No newline at end of file diff --git a/library/flux_train_utils.py b/library/flux_train_utils.py index 5e0eafc..6a56811 100644 --- a/library/flux_train_utils.py +++ b/library/flux_train_utils.py @@ -13,7 +13,8 @@ from transformers import CLIPTextModel from tqdm import tqdm from PIL import Image -from . import flux_models, flux_utils, strategy_base +from safetensors.torch import save_file +from . import flux_models, flux_utils, strategy_base, train_util from .sd3_train_utils import load_prompts from .device_utils import init_ipex, clean_memory_on_device @@ -182,7 +183,6 @@ def sample_image_inference( # sample image weight_dtype = ae.dtype # TOFO give dtype as argument - print("WEIGHT DTYPE: ", weight_dtype) packed_latent_height = height // 16 packed_latent_width = width // 16 noise = torch.randn( @@ -194,7 +194,7 @@ def sample_image_inference( generator=torch.Generator(device=accelerator.device).manual_seed(seed) if seed is not None else None, ) timesteps = get_schedule(sample_steps, noise.shape[1], shift=True) # FLUX.1 dev -> shift=True - print("TIMESTEPS: ", timesteps) + #print("TIMESTEPS: ", timesteps) img_ids = flux_utils.prepare_img_ids(1, packed_latent_height, packed_latent_width).to(accelerator.device, weight_dtype) with accelerator.autocast(), torch.no_grad(): @@ -280,10 +280,7 @@ def denoise( guidance: float = 4.0, ): # this is ignored for schnell - print("TRANSFORMER DTYPE: ", model.dtype) - print("IMAGE DTYPE: ", img.dtype) guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype) - print("GUIDANCE VECTOR: ", guidance_vec) comfy_pbar = ProgressBar(total=len(timesteps)) for t_curr, t_prev in zip(tqdm(timesteps[:-1]), timesteps[1:]): t_vec = torch.full((img.shape[0],), t_curr, dtype=img.dtype, device=img.device) @@ -292,4 +289,263 @@ def denoise( img = img + (t_prev - t_curr) * pred comfy_pbar.update(1) - return img \ No newline at end of file + return img + +# endregion + + +# region train +def get_sigmas(noise_scheduler, timesteps, device, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(device) + timesteps = timesteps.to(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 + + +def get_noisy_model_input_and_timesteps( + args, noise_scheduler, latents, noise, device, dtype +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + bsz = latents.shape[0] + sigmas = None + + if args.timestep_sampling == "uniform" or args.timestep_sampling == "sigmoid": + # Simple random t-based noise sampling + if args.timestep_sampling == "sigmoid": + # https://github.com/XLabs-AI/x-flux/tree/main + t = torch.sigmoid(args.sigmoid_scale * torch.randn((bsz,), device=device)) + else: + t = torch.rand((bsz,), device=device) + timesteps = t * 1000.0 + t = t.view(-1, 1, 1, 1) + noisy_model_input = (1 - t) * latents + t * noise + else: + # Sample a random timestep for each image + # for weighting schemes where we sample timesteps non-uniformly + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + timesteps = noise_scheduler.timesteps[indices].to(device=device) + + # Add noise according to flow matching. + sigmas = get_sigmas(noise_scheduler, timesteps, device, n_dim=latents.ndim, dtype=dtype) + noisy_model_input = sigmas * noise + (1.0 - sigmas) * latents + + return noisy_model_input, timesteps, sigmas + + +def apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas): + weighting = None + if args.model_prediction_type == "raw": + pass + elif args.model_prediction_type == "additive": + # add the model_pred to the noisy_model_input + model_pred = model_pred + noisy_model_input + elif args.model_prediction_type == "sigma_scaled": + # apply sigma scaling + model_pred = model_pred * (-sigmas) + noisy_model_input + + # these weighting schemes use a uniform timestep sampling + # and instead post-weight the loss + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + + return model_pred, weighting + + +def save_models(ckpt_path: str, flux: flux_models.Flux, sai_metadata: Optional[dict], save_dtype: Optional[torch.dtype] = None): + state_dict = {} + + def update_sd(prefix, sd): + for k, v in sd.items(): + key = prefix + k + if save_dtype is not None: + v = v.detach().clone().to("cpu").to(save_dtype) + state_dict[key] = v + + update_sd("", flux.state_dict()) + + save_file(state_dict, ckpt_path, metadata=sai_metadata) + + +def save_flux_model_on_train_end( + args: argparse.Namespace, save_dtype: torch.dtype, epoch: int, global_step: int, flux: flux_models.Flux +): + def sd_saver(ckpt_file, epoch_no, global_step): + sai_metadata = train_util.get_sai_model_spec(None, args, False, False, False, is_stable_diffusion_ckpt=True, flux="dev") + save_models(ckpt_file, flux, sai_metadata, save_dtype) + + train_util.save_sd_model_on_train_end_common(args, True, True, epoch, global_step, sd_saver, None) + + +# epochとstepの保存、メタデータにepoch/stepが含まれ引数が同じになるため、統合している +# on_epoch_end: Trueならepoch終了時、Falseならstep経過時 +def save_flux_model_on_epoch_end_or_stepwise( + args: argparse.Namespace, + on_epoch_end: bool, + accelerator, + save_dtype: torch.dtype, + epoch: int, + num_train_epochs: int, + global_step: int, + flux: flux_models.Flux, +): + def sd_saver(ckpt_file, epoch_no, global_step): + sai_metadata = train_util.get_sai_model_spec(None, args, False, False, False, is_stable_diffusion_ckpt=True, flux="dev") + save_models(ckpt_file, flux, sai_metadata, save_dtype) + + train_util.save_sd_model_on_epoch_end_or_stepwise_common( + args, + on_epoch_end, + accelerator, + True, + True, + epoch, + num_train_epochs, + global_step, + sd_saver, + None, + ) + + +# endregion + + +def add_flux_train_arguments(parser: argparse.ArgumentParser): + parser.add_argument( + "--clip_l", + type=str, + help="path to clip_l (*.sft or *.safetensors), should be float16 / clip_lのパス(*.sftまたは*.safetensors)、float16が前提", + ) + parser.add_argument( + "--t5xxl", + type=str, + help="path to t5xxl (*.sft or *.safetensors), should be float16 / t5xxlのパス(*.sftまたは*.safetensors)、float16が前提", + ) + parser.add_argument("--ae", type=str, help="path to ae (*.sft or *.safetensors) / aeのパス(*.sftまたは*.safetensors)") + parser.add_argument( + "--t5xxl_max_token_length", + type=int, + default=None, + help="maximum token length for T5-XXL. if omitted, 256 for schnell and 512 for dev" + " / T5-XXLの最大トークン長。省略された場合、schnellの場合は256、devの場合は512", + ) + parser.add_argument( + "--apply_t5_attn_mask", + action="store_true", + help="apply attention mask (zero embs) to T5-XXL / T5-XXLにアテンションマスク(ゼロ埋め)を適用する", + ) + parser.add_argument( + "--cache_text_encoder_outputs", action="store_true", help="cache text encoder outputs / text encoderの出力をキャッシュする" + ) + parser.add_argument( + "--cache_text_encoder_outputs_to_disk", + action="store_true", + help="cache text encoder outputs to disk / text encoderの出力をディスクにキャッシュする", + ) + parser.add_argument( + "--text_encoder_batch_size", + type=int, + default=None, + help="text encoder batch size (default: None, use dataset's batch size)" + + " / text encoderのバッチサイズ(デフォルト: None, データセットのバッチサイズを使用)", + ) + parser.add_argument( + "--disable_mmap_load_safetensors", + action="store_true", + help="disable mmap load for safetensors. Speed up model loading in WSL environment / safetensorsのmmapロードを無効にする。WSL環境等でモデル読み込みを高速化できる", + ) + + # copy from Diffusers + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument("--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme.") + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=3.5, + help="the FLUX.1 dev variant is a guidance distilled model", + ) + + parser.add_argument( + "--timestep_sampling", + choices=["sigma", "uniform", "sigmoid"], + default="sigma", + help="Method to sample timesteps: sigma-based, uniform random, or sigmoid of random normal. / タイムステップをサンプリングする方法:sigma、random uniform、またはrandom normalのsigmoid。", + ) + parser.add_argument( + "--sigmoid_scale", + type=float, + default=1.0, + help='Scale factor for sigmoid timestep sampling (only used when timestep-sampling is "sigmoid"). / sigmoidタイムステップサンプリングの倍率(timestep-samplingが"sigmoid"の場合のみ有効)。', + ) + parser.add_argument( + "--model_prediction_type", + choices=["raw", "additive", "sigma_scaled"], + default="sigma_scaled", + help="How to interpret and process the model prediction: " + "raw (use as is), additive (add to noisy input), sigma_scaled (apply sigma scaling)." + " / モデル予測の解釈と処理方法:" + "raw(そのまま使用)、additive(ノイズ入力に加算)、sigma_scaled(シグマスケーリングを適用)。", + ) + parser.add_argument( + "--discrete_flow_shift", + type=float, + default=3.0, + help="Discrete flow shift for the Euler Discrete Scheduler, default is 3.0. / Euler Discrete Schedulerの離散フローシフト、デフォルトは3.0。", + ) \ No newline at end of file diff --git a/library/train_util.py b/library/train_util.py index 6f1b4c6..301d06d 100644 --- a/library/train_util.py +++ b/library/train_util.py @@ -2631,7 +2631,7 @@ class MinimalDataset(BaseDataset): raise NotImplementedError -def load_arbitrary_dataset(args, tokenizer) -> MinimalDataset: +def load_arbitrary_dataset(args, tokenizer=None) -> MinimalDataset: module = ".".join(args.dataset_class.split(".")[:-1]) dataset_class = args.dataset_class.split(".")[-1] module = importlib.import_module(module) diff --git a/nodes.py b/nodes.py index cb85ec8..d881f89 100644 --- a/nodes.py +++ b/nodes.py @@ -12,12 +12,14 @@ from pathlib import Path script_directory = os.path.dirname(os.path.abspath(__file__)) from .flux_train_network_comfy import FluxNetworkTrainer - +from .library import flux_train_utils as flux_train_utils +from .flux_train_comfy import FluxTrainer +from .flux_train_comfy import setup_parser as train_setup_parser from .library.device_utils import init_ipex init_ipex() from .library import train_util -from .train_network import setup_parser +from .train_network import setup_parser as train_network_setup_parser import logging logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') @@ -59,15 +61,15 @@ class TrainDatasetConfig: @classmethod def INPUT_TYPES(s): return {"required": { - "width": ("INT",{"min": 64, "default": 512}), - "height": ("INT",{"min": 64, "default": 512}), + "width": ("INT",{"min": 64, "default": 1024, "tooltip": "image width when bucketing is not used, also the default validation sampling width"}), + "height": ("INT",{"min": 64, "default": 1024, "tooltip": "image height when bucketing is not used, also the default validation sampling height"}), "batch_size": ("INT",{"min": 1, "default": 2, "tooltip": "Higher batch size uses more memory and generalizes the training more. "}), - "dataset_path": ("STRING",{"multiline": True, "default": ""}), + "dataset_path": ("STRING",{"multiline": True, "default": "", "tooltip": "path to dataset, root is ComfyUI folder"}), "class_tokens": ("STRING",{"multiline": True, "default": ""}), "enable_bucket": ("BOOLEAN",{"default": True, "tooltip": "enable buckets for multi aspect ratio training"}), "bucket_no_upscale": ("BOOLEAN",{"default": False, "tooltip": "bucket reso is defined by image size automatically"}), "min_bucket_reso": ("INT",{"min": 64, "default": 256}), - "max_bucket_resos": ("STRING",{"default": "1024, 768, 512"}), + "max_bucket_resos": ("STRING",{"default": "1024, 768, 512", "tooltip": "comma separated list of bucket resos, when multiple are given the nearest to the original is used"}), "color_aug": ("BOOLEAN",{"default": False, "tooltip": "enable weak color augmentation"}), "flip_aug": ("BOOLEAN",{"default": False, "tooltip": "enable horizontal flip augmentation"}), "dataset_repeats": ("INT", {"default": 1, "min": 1, "tooltip": "number of times to repeat dataset for an epoch"}), @@ -113,20 +115,43 @@ class TrainDatasetConfig: "dataset": toml.dumps(dataset) } return (dataset_settings,) + +class OptimizerConfig: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "optimizer_type": (["adamw8bit", "adafactor", "prodigy"], {"default": "adamw8bit", "tooltip": "optimizer type"}), + "max_grad_norm": ("FLOAT",{"default": 1.0, "min": 0.0, "tooltip": "gradient clipping"}), + "lr_scheduler": (["constant", "cosine", "cosine_with_restarts", "polynomial", "constant_with_warmup", "adafactor"], {"default": "constant", "tooltip": "learning rate scheduler"}), + "lr_warmup_steps": ("INT",{"default": 0, "min": 0, "tooltip": "learning rate warmup steps"}), + "lr_scheduler_num_cycles": ("INT",{"default": 1, "min": 1, "tooltip": "learning rate scheduler num cycles"}), + "lr_scheduler_power": ("FLOAT",{"default": 1.0, "min": 0.0, "tooltip": "learning rate scheduler power"}), + }, + } + + RETURN_TYPES = ("ARGS",) + RETURN_NAMES = ("optimizer_settings",) + FUNCTION = "create_config" + CATEGORY = "FluxTrainer" + + def create_config(self, **kwargs): + + return (kwargs,) -class InitFluxTraining: + +class InitFluxLoRATraining: @classmethod def INPUT_TYPES(s): return {"required": { "flux_models": ("TRAIN_FLUX_MODELS",), "dataset_settings": ("TOML_DATASET",), + "optimizer_settings": ("ARGS",), "output_name": ("STRING", {"default": "flux_lora", "multiline": False}), "output_dir": ("STRING", {"default": "flux_trainer_output", "multiline": False}), "network_dim": ("INT", {"default": 4, "min": 1, "max": 256, "step": 1, "tooltip": "network dim"}), "learning_rate": ("FLOAT", {"default": 4e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "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"}), - "optimizer_type": (["adamw8bit", "adafactor", "prodigy"], {"default": "adamw8bit", "tooltip": "optimizer type"}), "max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 10000, "step": 1, "tooltip": "max number of training steps"}), "network_train_unet_only": ("BOOLEAN", {"default": True, "tooltip": "wheter to train the text encoder"}), "text_encoder_lr": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}), @@ -135,13 +160,14 @@ class InitFluxTraining: "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"}), - "weighting_scheme": (["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"],), + "weighting_scheme": (["logit_normal", "sigma_sqrt", "mode", "cosmap", "none"],), "logit_mean": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "mean to use when using the logit_normal weighting scheme"}), "logit_std": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01,"tooltip": "std to use when using the logit_normal weighting scheme"}), "mode_scale": ("FLOAT", {"default": 1.29, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Scale of mode weighting scheme. Only effective when using the mode as the weighting_scheme"}), "timestep_sampling": (["sigmoid", "uniform", "sigma"], {"tooltip": "method to sample timestep"}), "sigmoid_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for sigmoid timestep sampling (only used when timestep-sampling is sigmoid"}), "model_prediction_type": (["raw", "additive", "sigma_scaled"], {"tooltip": "How to interpret and process the model prediction: raw (use as is), additive (add to noisy input), sigma_scaled (apply sigma scaling)."}), + "guidance_scale": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 32.0, "step": 0.01, "tooltip": "guidance scale, for Flux training should be 1.0"}), "discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}), "highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}), "fp8_base": ("BOOLEAN", {"default": True, "tooltip": "use fp8 for base model"}), @@ -157,13 +183,13 @@ class InitFluxTraining: FUNCTION = "init_training" CATEGORY = "FluxTrainer" - def init_training(self, flux_models, dataset_settings, sample_prompts, output_name, optimizer_type, attention_mode, training_dtype, save_dtype, **kwargs,): + def init_training(self, flux_models, dataset_settings, optimizer_settings, sample_prompts, output_name, attention_mode, training_dtype, save_dtype, **kwargs,): mm.soft_empty_cache() dataset = dataset_settings["dataset"] dataset_repeats = dataset_settings["repeats"] - parser = setup_parser() + parser = train_network_setup_parser() args, _ = parser.parse_known_args() if kwargs.get("cache_latents") == "memory": @@ -219,8 +245,6 @@ class InitFluxTraining: "output_dir": output_dir, "output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}", "loss_type": "l2", - "optimizer_type": optimizer_type, - "guidance_scale": 3.5, "width" : int(width), "height" : int(height), } @@ -236,13 +260,14 @@ class InitFluxTraining: } config_dict.update(training_dtype_settings.get(training_dtype, {})) - if optimizer_type == "adafactor": + if optimizer_settings["optimizer_type"] == "adafactor": config_dict["optimizer_args"] = [ "relative_step=False", "scale_parameter=False", "warmup_init=False" ] config_dict.update(kwargs) + config_dict.update(optimizer_settings) for key, value in config_dict.items(): setattr(args, key, value) @@ -251,7 +276,7 @@ class InitFluxTraining: network_trainer = FluxNetworkTrainer() training_loop = network_trainer.init_train(args) - final_output_lora_path = os.path.join(output_dir, "output", output_name) + final_output_lora_path = os.path.join(output_dir, output_name) epochs_count = network_trainer.num_train_epochs @@ -260,13 +285,161 @@ class InitFluxTraining: "training_loop": training_loop, } return (trainer, epochs_count, final_output_lora_path) + +class InitFluxTraining: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "flux_models": ("TRAIN_FLUX_MODELS",), + "dataset_settings": ("TOML_DATASET",), + "optimizer_settings": ("OPTIMIZER_SETTINGS",), + "output_name": ("STRING", {"default": "flux", "multiline": False}), + "output_dir": ("STRING", {"default": "flux_trainer_output", "multiline": False, "tooltip": "output directory, root is ComfyUI folder"}), + "learning_rate": ("FLOAT", {"default": 5e-5, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "learning rate"}), + "max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 10000, "step": 1, "tooltip": "max number of training steps"}), + "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"}), + "weighting_scheme": (["logit_normal", "sigma_sqrt", "mode", "cosmap", "none"],), + "logit_mean": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "mean to use when using the logit_normal weighting scheme"}), + "logit_std": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01,"tooltip": "std to use when using the logit_normal weighting scheme"}), + "mode_scale": ("FLOAT", {"default": 1.29, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Scale of mode weighting scheme. Only effective when using the mode as the weighting_scheme"}), + "loss_type": (["l1", "l2", "huber", "smooth_l1"], {"default": "l2", "tooltip": "loss type"}), + "timestep_sampling": (["sigmoid", "uniform", "sigma"], {"tooltip": "method to sample timestep"}), + "sigmoid_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for sigmoid timestep sampling (only used when timestep-sampling is sigmoid"}), + "model_prediction_type": (["raw", "additive", "sigma_scaled"], {"tooltip": "How to interpret and process the model prediction: raw (use as is), additive (add to noisy input), sigma_scaled (apply sigma scaling)."}), + "cpu_offload_checkpointing": ("BOOLEAN", {"default": True, "tooltip": "offload the gradient checkpointing to CPU. This reduces VRAM usage for about 2GB"}), + "blockwise_fused_optimizer": ("BOOLEAN", {"default": True, "tooltip": "enables the fusing of the optimizer for each block"}), + "single_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "number of single blocks to swap. The default is 0. This option must be combined with blockwise_fused_optimizer"}), + "double_blocks_to_swap": ("INT", {"default": 6, "min": 0, "max": 100, "step": 1, "tooltip": "number of double blocks to swap. This option must be combined with blockwise_fused_optimizer"}), + "guidance_scale": ("FLOAT", {"default": 3.5, "min": 1.0, "max": 32.0, "step": 0.01, "tooltip": "guidance scale"}), + "discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}), + "highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}), + "fp8_base": ("BOOLEAN", {"default": False, "tooltip": "use fp8 for base model"}), + "full_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "to use the full fp16/bf16 training"}), + "save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn"], {"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 `|`"}), + }, + } + + RETURN_TYPES = ("NETWORKTRAINER", "INT", "STRING", ) + RETURN_NAMES = ("network_trainer", "epochs_count", "output_path",) + FUNCTION = "init_training" + CATEGORY = "FluxTrainer" + + def init_training(self, flux_models, optimizer_settings, dataset_settings, sample_prompts, output_name, optimizer_type, + attention_mode, full_dtype, save_dtype, **kwargs,): + mm.soft_empty_cache() + + dataset = dataset_settings["dataset"] + dataset_repeats = dataset_settings["repeats"] + + parser = train_setup_parser() + args, _ = parser.parse_known_args() + + if kwargs.get("cache_latents") == "memory": + kwargs["cache_latents"] = True + kwargs["cache_latents_to_disk"] = False + elif kwargs.get("cache_latents") == "disk": + kwargs["cache_latents"] = True + kwargs["cache_latents_to_disk"] = True + kwargs["caption_dropout_rate"] = 0.0 + kwargs["shuffle_caption"] = False + kwargs["token_warmup_step"] = 0.0 + kwargs["caption_tag_dropout_rate"] = 0.0 + else: + kwargs["cache_latents"] = False + kwargs["cache_latents_to_disk"] = False + + if kwargs.get("cache_text_encoder_outputs") == "memory": + kwargs["cache_text_encoder_outputs"] = True + kwargs["cache_text_encoder_outputs_to_disk"] = False + elif kwargs.get("cache_text_encoder_outputs") == "disk": + kwargs["cache_text_encoder_outputs"] = True + kwargs["cache_text_encoder_outputs_to_disk"] = True + else: + kwargs["cache_text_encoder_outputs"] = False + kwargs["cache_text_encoder_outputs_to_disk"] = False + + + output_dir = os.path.join(script_directory, "output") + if '|' in sample_prompts: + prompts = sample_prompts.split('|') + else: + prompts = [sample_prompts] + + width, height = toml.loads(dataset)["datasets"][0]["resolution"] + config_dict = { + "sample_prompts": prompts, + "save_precision": save_dtype, + "dataset_repeats": dataset_repeats, + "mixed_precision": "bf16", + "num_cpu_threads_per_process": 1, + "pretrained_model_name_or_path": flux_models["transformer"], + "clip_l": flux_models["clip_l"], + "t5xxl": flux_models["t5"], + "ae": flux_models["vae"], + "save_model_as": "safetensors", + "persistent_data_loader_workers": False, + "max_data_loader_n_workers": 0, + "seed": 42, + "gradient_checkpointing": True, + "save_precision": "bf16", + "dataset_config": dataset, + "output_dir": output_dir, + "output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}", + "optimizer_type": optimizer_type, + "width" : int(width), + "height" : int(height), + + } + attention_settings = { + "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True}, + "xformers": {"mem_eff_attn": True, "xformers": True, "spda": False} + } + config_dict.update(attention_settings.get(attention_mode, {})) + + full_dtype_settings = { + "fp16": {"full_fp16": True, "full_bf16": False}, + "bf16": {"full_bf16": True, "full_fp16": False} + } + config_dict.update(full_dtype_settings.get(full_dtype, {})) + + if optimizer_settings["optimizer_type"] == "adafactor": + config_dict["optimizer_args"] = [ + "relative_step=False", + "scale_parameter=False", + "warmup_init=False" + ] + config_dict["max_grad_norm"] = 0 + config_dict.update(kwargs) + config_dict.update(optimizer_settings) + + for key, value in config_dict.items(): + setattr(args, key, value) + + with torch.inference_mode(False): + network_trainer = FluxTrainer() + training_loop = network_trainer.init_train(args) + + final_output_path = os.path.join(output_dir, output_name) + + epochs_count = network_trainer.num_train_epochs + + trainer = { + "network_trainer": network_trainer, + "training_loop": training_loop, + } + return (trainer, epochs_count, final_output_path) class FluxTrainLoop: @classmethod def INPUT_TYPES(s): return {"required": { "network_trainer": ("NETWORKTRAINER",), - "steps": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1}), + "steps": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1, "tooltip": "the step point in training to validate/save"}), }, } @@ -309,29 +482,30 @@ class FluxTrainSave: }, } - RETURN_TYPES = ("NETWORKTRAINER", "STRING",) - RETURN_NAMES = ("network_trainer","lora_path",) + RETURN_TYPES = ("NETWORKTRAINER", "STRING", "INT",) + RETURN_NAMES = ("network_trainer","lora_path", "steps",) FUNCTION = "endtrain" CATEGORY = "FluxTrainer" def endtrain(self, network_trainer, save_state): with torch.inference_mode(False): trainer = network_trainer["network_trainer"] + global_step = trainer.global_step - ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, trainer.global_step) - trainer.save_model(ckpt_name, trainer.accelerator.unwrap_model(trainer.network), trainer.global_step, trainer.current_epoch.value + 1) + ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, global_step) + trainer.save_model(ckpt_name, trainer.accelerator.unwrap_model(trainer.network), global_step, trainer.current_epoch.value + 1) - remove_step_no = train_util.get_remove_step_no(trainer.args, trainer.global_step) + remove_step_no = train_util.get_remove_step_no(trainer.args, global_step) if remove_step_no is not None: remove_ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, remove_step_no) trainer.remove_model(remove_ckpt_name) if save_state: - train_util.save_and_remove_state_stepwise(trainer.args, trainer.accelerator, trainer.global_step) + train_util.save_and_remove_state_stepwise(trainer.args, trainer.accelerator, global_step) - lora_path = os.path.join(trainer.args.output_dir, "output", ckpt_name) + lora_path = os.path.join(trainer.args.output_dir, ckpt_name) - return (network_trainer, lora_path) + return (network_trainer, lora_path, global_step) class FluxTrainEnd: @classmethod @@ -366,7 +540,7 @@ class FluxTrainEnd: network_trainer.save_model(ckpt_name, network, network_trainer.global_step, network_trainer.num_train_epochs, force_sync_upload=True) logger.info("model saved.") - final_output_lora_path = os.path.join(network_trainer.args.output_dir, "output", network_trainer.args.output_name) + final_output_lora_path = os.path.join(network_trainer.args.output_dir, network_trainer.args.output_name) # metadata metadata = json.dumps(network_trainer.metadata, indent=2) @@ -421,16 +595,22 @@ class FluxTrainValidate: training_loop = network_trainer["training_loop"] network_trainer = network_trainer["network_trainer"] - image_tensors = network_trainer.sample_images( + params = ( network_trainer.accelerator, network_trainer.args, network_trainer.current_epoch.value, network_trainer.global_step, + network_trainer.unet, network_trainer.vae, network_trainer.text_encoder, - network_trainer.unet, + network_trainer.sample_prompts_te_outputs, validation_settings - ) + ) + + if not network_trainer.args.split_mode: + image_tensors = flux_train_utils.sample_images(*params) + else: + image_tensors = network_trainer.sample_images_split_mode(*params) trainer = { "network_trainer": network_trainer, @@ -749,8 +929,7 @@ class UploadToHuggingFace: "network_trainer": ("NETWORKTRAINER",), "source_path": ("STRING", {"default": ""}), "repo_id": ("STRING",{"default": ""}), - "path_in_repo": ("STRING",{"default": "model"}), - "revision": ("STRING", {"default": "main"}), + "revision": ("STRING", {"default": ""}), "private": ("BOOLEAN", {"default": True, "tooltip": "If creating a new repo, leave it private"}), }, "optional": { @@ -758,76 +937,89 @@ class UploadToHuggingFace: } } - RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("status",) + RETURN_TYPES = ("NETWORKTRAINER", "STRING",) + RETURN_NAMES = ("network_trainer","status",) FUNCTION = "upload" CATEGORY = "FluxTrainer" - def upload(self, source_path, network_trainer, repo_id, path_in_repo, private, revision,token): - from huggingface_hub import HfApi - - with open(os.path.join(script_directory, "hf_token.json"), "r") as file: - token_data = json.load(file) - token = token_data["hf_token"] - - # Save metadata to a JSON file - metadata = network_trainer["network_trainer"].metadata - metadata_file_path = Path(source_path) / "metadata.json" - with open(metadata_file_path, 'w') as f: - json.dump(metadata, f) - - repo_type = "model" - api = HfApi(token=token) - - try: - api.repo_info(repo_id=repo_id, revision=revision, repo_type=repo_type) - repo_exists = True - except: - repo_exists = False - - if not repo_exists(repo_id=repo_id, repo_type=repo_type, token=token): - try: - api.create_repo(repo_id=repo_id, repo_type=repo_type, private=private) - except Exception as e: # Checked for RepositoryNotFoundError, but other exceptions could be problematic - logger.error("===========================================") - logger.error(f"failed to create HuggingFace repo: {e}") - logger.error("===========================================") - - is_folder = (type(source_path) == str and os.path.isdir(source_path)) or (isinstance(source_path, Path) and source_path.is_dir()) - - try: - if is_folder: - api.upload_folder( - repo_id=repo_id, - repo_type=repo_type, - folder_path=source_path, - path_in_repo=path_in_repo, - ) - else: - api.upload_file( - repo_id=repo_id, - repo_type=repo_type, - path_or_fileobj=source_path, - path_in_repo=path_in_repo, - ) - # Upload the metadata file separately if it's not a folder upload - if not is_folder: - api.upload_file( - repo_id=repo_id, - repo_type=repo_type, - path_or_fileobj=str(metadata_file_path), - path_in_repo=path_in_repo + '/metadata.json', - ) - status = "Uploaded to HuggingFace succesfully" - except Exception as e: # RuntimeErrorを確認済みだが他にあると困るので - logger.error("===========================================") - logger.error(f"failed to upload to HuggingFace / HuggingFaceへのアップロードに失敗しました : {e}") - logger.error("===========================================") - status = f"Failed to upload to HuggingFace {e}" + def upload(self, source_path, network_trainer, repo_id, private, revision, token=""): + with torch.inference_mode(False): + from huggingface_hub import HfApi - return (status,) + if not token: + with open(os.path.join(script_directory, "hf_token.json"), "r") as file: + token_data = json.load(file) + token = token_data["hf_token"] + print(token) + + # Save metadata to a JSON file + directory_path = os.path.dirname(os.path.dirname(source_path)) + file_name = os.path.basename(source_path) + + metadata = network_trainer["network_trainer"].metadata + metadata_file_path = os.path.join(directory_path, "metadata.json") + with open(metadata_file_path, 'w') as f: + json.dump(metadata, f, indent=4) + + repo_type = None + api = HfApi(token=token) + + try: + api.repo_info( + repo_id=repo_id, + revision=revision if revision != "" else None, + repo_type=repo_type) + repo_exists = True + logger.info(f"Repository {repo_id} exists.") + except Exception as e: # Catching a more specific exception would be better if you know what to expect + repo_exists = False + logger.error(f"Repository {repo_id} does not exist. Exception: {e}") + + if not repo_exists: + try: + api.create_repo(repo_id=repo_id, repo_type=repo_type, private=private) + except Exception as e: # Checked for RepositoryNotFoundError, but other exceptions could be problematic + logger.error("===========================================") + logger.error(f"failed to create HuggingFace repo: {e}") + logger.error("===========================================") + + is_folder = (type(source_path) == str and os.path.isdir(source_path)) or (isinstance(source_path, Path) and source_path.is_dir()) + print(source_path, is_folder) + + try: + if is_folder: + api.upload_folder( + repo_id=repo_id, + repo_type=repo_type, + folder_path=source_path, + path_in_repo=file_name, + ) + else: + api.upload_file( + repo_id=repo_id, + repo_type=repo_type, + path_or_fileobj=source_path, + path_in_repo=file_name, + ) + # Upload the metadata file separately if it's not a folder upload + if not is_folder: + api.upload_file( + repo_id=repo_id, + repo_type=repo_type, + path_or_fileobj=str(metadata_file_path), + path_in_repo='metadata.json', + ) + status = "Uploaded to HuggingFace succesfully" + except Exception as e: # RuntimeErrorを確認済みだが他にあると困るので + logger.error("===========================================") + logger.error(f"failed to upload to HuggingFace / HuggingFaceへのアップロードに失敗しました : {e}") + logger.error("===========================================") + status = f"Failed to upload to HuggingFace {e}" + + return (network_trainer, status,) NODE_CLASS_MAPPINGS = { + "InitFluxLoRATraining": InitFluxLoRATraining, "InitFluxTraining": InitFluxTraining, "FluxTrainModelSelect": FluxTrainModelSelect, "TrainDatasetConfig": TrainDatasetConfig, @@ -838,9 +1030,11 @@ NODE_CLASS_MAPPINGS = { "FluxTrainEnd": FluxTrainEnd, "FluxTrainSave": FluxTrainSave, "FluxKohyaInferenceSampler": FluxKohyaInferenceSampler, - "UploadToHuggingFace": UploadToHuggingFace + "UploadToHuggingFace": UploadToHuggingFace, + "OptimizerConfig": OptimizerConfig } NODE_DISPLAY_NAME_MAPPINGS = { + "InitFluxLoRATraining": "Init Flux LoRA Training", "InitFluxTraining": "Init Flux Training", "FluxTrainModelSelect": "FluxTrain ModelSelect", "TrainDatasetConfig": "Train Dataset Config", @@ -851,5 +1045,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FluxTrainEnd": "Flux Train End", "FluxTrainSave": "Flux Train Save", "FluxKohyaInferenceSampler": "Flux Kohya Inference Sampler", - "UploadToHuggingFace": "Upload To HuggingFace" + "UploadToHuggingFace": "Upload To HuggingFace", + "OptimizerConfig": "Optimizer Config" }