diff --git a/.gitignore b/.gitignore
index 5a877c2..546782a 100644
--- a/.gitignore
+++ b/.gitignore
@@ -18,3 +18,4 @@ scepter.egg-info
#MANIFEST.in
*resources
*.ipynb_checkpoints*
+*.vscode
diff --git a/asset/images/ace/logo.png b/asset/images/ace/logo.png
deleted file mode 100644
index d1b240e..0000000
Binary files a/asset/images/ace/logo.png and /dev/null differ
diff --git a/asset/images/ace/teaser_dy.gif b/asset/images/ace/teaser_dy.gif
deleted file mode 100644
index 4291dbf..0000000
Binary files a/asset/images/ace/teaser_dy.gif and /dev/null differ
diff --git a/asset/images/ace/text.png b/asset/images/ace/text.png
deleted file mode 100644
index 8c00258..0000000
Binary files a/asset/images/ace/text.png and /dev/null differ
diff --git a/asset/images/flux_tuner/flux_tuner_1_1.webp b/asset/images/flux_tuner/flux_tuner_1_1.webp
new file mode 100644
index 0000000..cb8d9ff
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_1_1.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_1_2.webp b/asset/images/flux_tuner/flux_tuner_1_2.webp
new file mode 100644
index 0000000..266d3ce
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_1_2.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_1_3.webp b/asset/images/flux_tuner/flux_tuner_1_3.webp
new file mode 100644
index 0000000..00fb745
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_1_3.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_2_1.webp b/asset/images/flux_tuner/flux_tuner_2_1.webp
new file mode 100644
index 0000000..298fffc
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_2_1.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_2_2.webp b/asset/images/flux_tuner/flux_tuner_2_2.webp
new file mode 100644
index 0000000..6826a7a
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_2_2.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_2_3.webp b/asset/images/flux_tuner/flux_tuner_2_3.webp
new file mode 100644
index 0000000..68c65c0
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_2_3.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_3_1.webp b/asset/images/flux_tuner/flux_tuner_3_1.webp
new file mode 100644
index 0000000..538b9aa
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_3_1.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_3_2.webp b/asset/images/flux_tuner/flux_tuner_3_2.webp
new file mode 100644
index 0000000..2ee23a6
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_3_2.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_3_3.webp b/asset/images/flux_tuner/flux_tuner_3_3.webp
new file mode 100644
index 0000000..610f2e5
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_3_3.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_4_1.webp b/asset/images/flux_tuner/flux_tuner_4_1.webp
new file mode 100644
index 0000000..736a0f9
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_4_1.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_4_2.webp b/asset/images/flux_tuner/flux_tuner_4_2.webp
new file mode 100644
index 0000000..9d7b90c
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_4_2.webp differ
diff --git a/asset/images/flux_tuner/flux_tuner_4_3.webp b/asset/images/flux_tuner/flux_tuner_4_3.webp
new file mode 100644
index 0000000..3371937
Binary files /dev/null and b/asset/images/flux_tuner/flux_tuner_4_3.webp differ
diff --git a/asset/workflow/sdxl_base.jpg b/asset/workflow/sdxl_base.jpg
new file mode 100644
index 0000000..b004e30
Binary files /dev/null and b/asset/workflow/sdxl_base.jpg differ
diff --git a/asset/workflow/sdxl_base.json b/asset/workflow/sdxl_base.json
new file mode 100644
index 0000000..35c4d94
--- /dev/null
+++ b/asset/workflow/sdxl_base.json
@@ -0,0 +1,163 @@
+{
+ "last_node_id": 4,
+ "last_link_id": 3,
+ "nodes": [
+ {
+ "id": 2,
+ "type": "PreviewImage",
+ "pos": {
+ "0": 937,
+ "1": 243
+ },
+ "size": {
+ "0": 284.03680419921875,
+ "1": 246
+ },
+ "flags": {
+ "collapsed": false
+ },
+ "order": 2,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": 1
+ }
+ ],
+ "outputs": [],
+ "properties": {
+ "Node name for S&R": "PreviewImage"
+ },
+ "widgets_values": []
+ },
+ {
+ "id": 4,
+ "type": "ParameterNode",
+ "pos": {
+ "0": 0,
+ "1": 244
+ },
+ "size": {
+ "0": 311.9849853515625,
+ "1": 236.4765625
+ },
+ "flags": {},
+ "order": 0,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 3
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ParameterNode"
+ },
+ "widgets_values": [
+ "ddim",
+ 50,
+ 5,
+ 0.5,
+ "trailing",
+ 1024,
+ 1024,
+ 2024
+ ]
+ },
+ {
+ "id": 1,
+ "type": "ModelNode",
+ "pos": {
+ "0": 430,
+ "1": 243
+ },
+ "size": {
+ "0": 400,
+ "1": 234
+ },
+ "flags": {},
+ "order": 1,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "parameters",
+ "type": "CONDITIONING",
+ "link": 3,
+ "shape": 7
+ },
+ {
+ "name": "mantras",
+ "type": "CONDITIONING",
+ "link": null,
+ "shape": 7
+ },
+ {
+ "name": "tuners",
+ "type": "CONDITIONING",
+ "link": null,
+ "shape": 7
+ },
+ {
+ "name": "controls",
+ "type": "CONDITIONING",
+ "link": null,
+ "shape": 7
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 1
+ ],
+ "slot_index": 0
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ModelNode"
+ },
+ "widgets_values": [
+ "SD_XL1.0",
+ "ModelScope",
+ "a cat",
+ ""
+ ]
+ }
+ ],
+ "links": [
+ [
+ 1,
+ 1,
+ 0,
+ 2,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 3,
+ 4,
+ 0,
+ 1,
+ 0,
+ "CONDITIONING"
+ ]
+ ],
+ "groups": [],
+ "config": {},
+ "extra": {
+ "ds": {
+ "scale": 0.9229599817706423,
+ "offset": [
+ 128.91466506592707,
+ 43.07074577975898
+ ]
+ }
+ },
+ "version": 0.4
+}
\ No newline at end of file
diff --git a/asset/workflow/sdxl_base_mantra.jpg b/asset/workflow/sdxl_base_mantra.jpg
new file mode 100644
index 0000000..aa06a0c
Binary files /dev/null and b/asset/workflow/sdxl_base_mantra.jpg differ
diff --git a/asset/workflow/sdxl_base_mantra.json b/asset/workflow/sdxl_base_mantra.json
new file mode 100644
index 0000000..4a9360e
--- /dev/null
+++ b/asset/workflow/sdxl_base_mantra.json
@@ -0,0 +1,202 @@
+{
+ "last_node_id": 5,
+ "last_link_id": 4,
+ "nodes": [
+ {
+ "id": 2,
+ "type": "PreviewImage",
+ "pos": {
+ "0": 937,
+ "1": 243
+ },
+ "size": {
+ "0": 284.03680419921875,
+ "1": 246
+ },
+ "flags": {
+ "collapsed": false
+ },
+ "order": 3,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": 1
+ }
+ ],
+ "outputs": [],
+ "properties": {
+ "Node name for S&R": "PreviewImage"
+ },
+ "widgets_values": []
+ },
+ {
+ "id": 4,
+ "type": "ParameterNode",
+ "pos": {
+ "0": 0,
+ "1": 244
+ },
+ "size": {
+ "0": 311.9849853515625,
+ "1": 236.4765625
+ },
+ "flags": {},
+ "order": 0,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 3
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ParameterNode"
+ },
+ "widgets_values": [
+ "ddim",
+ 50,
+ 5,
+ 0.5,
+ "trailing",
+ 1024,
+ 1024,
+ 2024
+ ]
+ },
+ {
+ "id": 1,
+ "type": "ModelNode",
+ "pos": {
+ "0": 430,
+ "1": 243
+ },
+ "size": {
+ "0": 400,
+ "1": 234
+ },
+ "flags": {},
+ "order": 2,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "parameters",
+ "type": "CONDITIONING",
+ "link": 3,
+ "shape": 7
+ },
+ {
+ "name": "mantras",
+ "type": "CONDITIONING",
+ "link": 4,
+ "shape": 7
+ },
+ {
+ "name": "tuners",
+ "type": "CONDITIONING",
+ "link": null,
+ "shape": 7
+ },
+ {
+ "name": "controls",
+ "type": "CONDITIONING",
+ "link": null,
+ "shape": 7
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 1
+ ],
+ "slot_index": 0
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ModelNode"
+ },
+ "widgets_values": [
+ "SD_XL1.0",
+ "ModelScope",
+ "a cat",
+ ""
+ ]
+ },
+ {
+ "id": 5,
+ "type": "MantrasNode",
+ "pos": {
+ "0": 3,
+ "1": 547
+ },
+ "size": {
+ "0": 315,
+ "1": 58
+ },
+ "flags": {},
+ "order": 1,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 4
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "MantrasNode"
+ },
+ "widgets_values": [
+ "Flat 2D Art"
+ ]
+ }
+ ],
+ "links": [
+ [
+ 1,
+ 1,
+ 0,
+ 2,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 3,
+ 4,
+ 0,
+ 1,
+ 0,
+ "CONDITIONING"
+ ],
+ [
+ 4,
+ 5,
+ 0,
+ 1,
+ 1,
+ "CONDITIONING"
+ ]
+ ],
+ "groups": [],
+ "config": {},
+ "extra": {
+ "ds": {
+ "scale": 1.0152559799477068,
+ "offset": [
+ 58.35814123566077,
+ -65.12749162217047
+ ]
+ }
+ },
+ "version": 0.4
+}
\ No newline at end of file
diff --git a/asset/workflow/sdxl_base_mantra_tuner.jpg b/asset/workflow/sdxl_base_mantra_tuner.jpg
new file mode 100644
index 0000000..4283ffe
Binary files /dev/null and b/asset/workflow/sdxl_base_mantra_tuner.jpg differ
diff --git a/asset/workflow/sdxl_base_mantra_tuner.json b/asset/workflow/sdxl_base_mantra_tuner.json
new file mode 100644
index 0000000..16f72ba
--- /dev/null
+++ b/asset/workflow/sdxl_base_mantra_tuner.json
@@ -0,0 +1,242 @@
+{
+ "last_node_id": 6,
+ "last_link_id": 5,
+ "nodes": [
+ {
+ "id": 2,
+ "type": "PreviewImage",
+ "pos": {
+ "0": 937,
+ "1": 243
+ },
+ "size": {
+ "0": 284.03680419921875,
+ "1": 246
+ },
+ "flags": {
+ "collapsed": false
+ },
+ "order": 4,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": 1
+ }
+ ],
+ "outputs": [],
+ "properties": {
+ "Node name for S&R": "PreviewImage"
+ },
+ "widgets_values": []
+ },
+ {
+ "id": 4,
+ "type": "ParameterNode",
+ "pos": {
+ "0": 0,
+ "1": 244
+ },
+ "size": {
+ "0": 311.9849853515625,
+ "1": 236.4765625
+ },
+ "flags": {},
+ "order": 0,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 3
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ParameterNode"
+ },
+ "widgets_values": [
+ "ddim",
+ 50,
+ 5,
+ 0.5,
+ "trailing",
+ 1024,
+ 1024,
+ 2024
+ ]
+ },
+ {
+ "id": 5,
+ "type": "MantrasNode",
+ "pos": {
+ "0": 3,
+ "1": 547
+ },
+ "size": {
+ "0": 315,
+ "1": 58
+ },
+ "flags": {},
+ "order": 1,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 4
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "MantrasNode"
+ },
+ "widgets_values": [
+ "Flat 2D Art"
+ ]
+ },
+ {
+ "id": 1,
+ "type": "ModelNode",
+ "pos": {
+ "0": 430,
+ "1": 243
+ },
+ "size": {
+ "0": 400,
+ "1": 234
+ },
+ "flags": {},
+ "order": 3,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "parameters",
+ "type": "CONDITIONING",
+ "link": 3,
+ "shape": 7
+ },
+ {
+ "name": "mantras",
+ "type": "CONDITIONING",
+ "link": 4,
+ "shape": 7
+ },
+ {
+ "name": "tuners",
+ "type": "CONDITIONING",
+ "link": 5,
+ "shape": 7
+ },
+ {
+ "name": "controls",
+ "type": "CONDITIONING",
+ "link": null,
+ "shape": 7
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 1
+ ],
+ "slot_index": 0
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ModelNode"
+ },
+ "widgets_values": [
+ "SD_XL1.0",
+ "ModelScope",
+ "a cat",
+ ""
+ ]
+ },
+ {
+ "id": 6,
+ "type": "TunerNode",
+ "pos": {
+ "0": 8,
+ "1": 680
+ },
+ "size": {
+ "0": 315,
+ "1": 82
+ },
+ "flags": {},
+ "order": 2,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 5
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "TunerNode"
+ },
+ "widgets_values": [
+ "SD_XL1.0_Pencil Sketch Drawing",
+ 1
+ ]
+ }
+ ],
+ "links": [
+ [
+ 1,
+ 1,
+ 0,
+ 2,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 3,
+ 4,
+ 0,
+ 1,
+ 0,
+ "CONDITIONING"
+ ],
+ [
+ 4,
+ 5,
+ 0,
+ 1,
+ 1,
+ "CONDITIONING"
+ ],
+ [
+ 5,
+ 6,
+ 0,
+ 1,
+ 2,
+ "CONDITIONING"
+ ]
+ ],
+ "groups": [],
+ "config": {},
+ "extra": {
+ "ds": {
+ "scale": 0.6934334949441366,
+ "offset": [
+ 450.0923948318967,
+ 30.783653239467935
+ ]
+ }
+ },
+ "version": 0.4
+}
\ No newline at end of file
diff --git a/asset/workflow/sdxl_base_mantra_tuner_control.jpg b/asset/workflow/sdxl_base_mantra_tuner_control.jpg
new file mode 100644
index 0000000..092f9d9
Binary files /dev/null and b/asset/workflow/sdxl_base_mantra_tuner_control.jpg differ
diff --git a/asset/workflow/sdxl_base_mantra_tuner_control.json b/asset/workflow/sdxl_base_mantra_tuner_control.json
new file mode 100644
index 0000000..8198a01
--- /dev/null
+++ b/asset/workflow/sdxl_base_mantra_tuner_control.json
@@ -0,0 +1,366 @@
+{
+ "last_node_id": 11,
+ "last_link_id": 9,
+ "nodes": [
+ {
+ "id": 1,
+ "type": "ModelNode",
+ "pos": {
+ "0": 439,
+ "1": 123
+ },
+ "size": {
+ "0": 400,
+ "1": 234
+ },
+ "flags": {},
+ "order": 6,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "parameters",
+ "type": "CONDITIONING",
+ "link": 3,
+ "shape": 7
+ },
+ {
+ "name": "mantras",
+ "type": "CONDITIONING",
+ "link": 4,
+ "shape": 7
+ },
+ {
+ "name": "tuners",
+ "type": "CONDITIONING",
+ "link": 5,
+ "shape": 7
+ },
+ {
+ "name": "controls",
+ "type": "CONDITIONING",
+ "link": 6,
+ "shape": 7
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 1
+ ],
+ "slot_index": 0
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ModelNode"
+ },
+ "widgets_values": [
+ "SD_XL1.0",
+ "ModelScope",
+ "a cat",
+ ""
+ ]
+ },
+ {
+ "id": 5,
+ "type": "MantrasNode",
+ "pos": {
+ "0": 13,
+ "1": 336
+ },
+ "size": {
+ "0": 315,
+ "1": 58
+ },
+ "flags": {},
+ "order": 0,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 4
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "MantrasNode"
+ },
+ "widgets_values": [
+ "Flat 2D Art"
+ ]
+ },
+ {
+ "id": 6,
+ "type": "TunerNode",
+ "pos": {
+ "0": 12,
+ "1": 455
+ },
+ "size": {
+ "0": 315,
+ "1": 82
+ },
+ "flags": {},
+ "order": 1,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 5
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "TunerNode"
+ },
+ "widgets_values": [
+ "SD_XL1.0_Pencil Sketch Drawing",
+ 1
+ ]
+ },
+ {
+ "id": 11,
+ "type": "NoteNode",
+ "pos": {
+ "0": 845,
+ "1": 506
+ },
+ "size": {
+ "0": 398.3059387207031,
+ "1": 210.83267211914062
+ },
+ "flags": {},
+ "order": 2,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [],
+ "properties": {
+ "Node name for S&R": "NoteNode"
+ },
+ "widgets_values": [
+ "This is a sample for quickly setting up ComfyUI using the Scepter open-source library:\n\n1) The first run involves downloading the models. By default, we will automatically pull models from ModelScope. The initial download may take some time, and you can also adjust the model source to change the model address.\n\n2) Currently, it supports various base models, basic settings for some inference hyperparameters, mantra settings, tuning model settings, and conditional generation node.\n\n3) In the future, we will gradually integrate interesting features to enrich the use cases.\n\n"
+ ]
+ },
+ {
+ "id": 8,
+ "type": "LoadImage",
+ "pos": {
+ "0": 471,
+ "1": 451
+ },
+ "size": {
+ "0": 315,
+ "1": 314
+ },
+ "flags": {},
+ "order": 3,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 7
+ ]
+ },
+ {
+ "name": "MASK",
+ "type": "MASK",
+ "links": null
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "LoadImage"
+ },
+ "widgets_values": [
+ "cat.jpg",
+ "image"
+ ]
+ },
+ {
+ "id": 7,
+ "type": "ControlNode",
+ "pos": {
+ "0": 15,
+ "1": 600
+ },
+ "size": {
+ "0": 330,
+ "1": 198
+ },
+ "flags": {},
+ "order": 5,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "source_image",
+ "type": "IMAGE",
+ "link": 7
+ }
+ ],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 6
+ ]
+ },
+ {
+ "name": "Control Image",
+ "type": "IMAGE",
+ "links": [],
+ "slot_index": 1
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ControlNode"
+ },
+ "widgets_values": [
+ "SD_XL1.0_color",
+ "Color",
+ "CenterCrop",
+ 1,
+ 1024,
+ 1024
+ ]
+ },
+ {
+ "id": 2,
+ "type": "PreviewImage",
+ "pos": {
+ "0": 959,
+ "1": 121
+ },
+ "size": {
+ "0": 284.03680419921875,
+ "1": 246
+ },
+ "flags": {
+ "collapsed": false
+ },
+ "order": 7,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": 1
+ }
+ ],
+ "outputs": [],
+ "properties": {
+ "Node name for S&R": "PreviewImage"
+ },
+ "widgets_values": []
+ },
+ {
+ "id": 4,
+ "type": "ParameterNode",
+ "pos": {
+ "0": 13,
+ "1": 43
+ },
+ "size": {
+ "0": 311.9849853515625,
+ "1": 236.4765625
+ },
+ "flags": {},
+ "order": 4,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "Result",
+ "type": "CONDITIONING",
+ "links": [
+ 3
+ ]
+ }
+ ],
+ "properties": {
+ "Node name for S&R": "ParameterNode"
+ },
+ "widgets_values": [
+ "ddim",
+ 50,
+ 5,
+ 0.5,
+ "trailing",
+ 1024,
+ 1024,
+ 2024
+ ]
+ }
+ ],
+ "links": [
+ [
+ 1,
+ 1,
+ 0,
+ 2,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 3,
+ 4,
+ 0,
+ 1,
+ 0,
+ "CONDITIONING"
+ ],
+ [
+ 4,
+ 5,
+ 0,
+ 1,
+ 1,
+ "CONDITIONING"
+ ],
+ [
+ 5,
+ 6,
+ 0,
+ 1,
+ 2,
+ "CONDITIONING"
+ ],
+ [
+ 6,
+ 7,
+ 0,
+ 1,
+ 3,
+ "CONDITIONING"
+ ],
+ [
+ 7,
+ 8,
+ 0,
+ 7,
+ 0,
+ "IMAGE"
+ ]
+ ],
+ "groups": [],
+ "config": {},
+ "extra": {
+ "ds": {
+ "scale": 0.7627768444385667,
+ "offset": [
+ 254.62860058494985,
+ 87.05734194144193
+ ]
+ }
+ },
+ "version": 0.4
+}
\ No newline at end of file
diff --git a/asset/workflow/workflow.jpg b/asset/workflow/workflow.jpg
new file mode 100644
index 0000000..f98f1b4
Binary files /dev/null and b/asset/workflow/workflow.jpg differ
diff --git a/readme.md b/readme.md
index f06838b..584fe90 100644
--- a/readme.md
+++ b/readme.md
@@ -18,7 +18,8 @@ SCEPTER offers 3 core components:
## 🎉 News
-- [2024.09]: We introduce **ACE**, an **A**ll-round **C**reator and **E**ditor adept at executing a diverse array of image editing tasks tailored to your specifications. Built upon the cutting-edge Diffusion Transformer architecture, ACE has been extensively trained on a comprehensive dataset to seamlessly interpret and execute any natural language instruction. For further information, please consult the [project page](https://ali-vilab.github.io/ace-page/).
+- [2024.10]: Support for inference and tuning with [FLUX](https://huggingface.co/black-forest-labs/FLUX.1-dev), as well as for building [ComfyUI](https://github.com/comfyanonymous/ComfyUI) workflows using this framework.
+- [🔥2024.09]: We introduce **ACE**, an **A**ll-round **C**reator and **E**ditor adept at executing a diverse array of image editing tasks tailored to your specifications. Built upon the cutting-edge Diffusion Transformer architecture, ACE has been extensively trained on a comprehensive dataset to seamlessly interpret and execute any natural language instruction. For further information, please consult the [project page](https://ali-vilab.github.io/ace-page/).
- [2024.07]: Support the inference and training of open-source generative models based on the [DiT](https://arxiv.org/abs/2212.09748) architecture, such as [SD3](https://arxiv.org/pdf/2403.03206) and [PixArt](https://arxiv.org/abs/2310.00426).
- [2024.05]: Introducing SCEPTER v1, supporting customized image edit tasks! Simply provide 10 image pairs, SCEPTER will tune an edit tuner for your own Image-to-Image tasks, like `Clay Style`, `De-Text`, `Segmentation`, etc.
- [2024.04]: New [StyleBooth](https://ali-vilab.github.io/stylebooth-page/) demo on SCEPTER Studio for`Text-Based Style Editing`.
@@ -32,15 +33,78 @@ SCEPTER offers 3 core components:
## 🖼 Gallery for Recent Works
-###
+###
ACE is a unified foundational model framework that supports a wide range of visual generation tasks. By defining CU for unifying multi-modal inputs across different tasks and incorporating long-
context CU, we introduce historical contextual information into visual generation tasks, paving
the way for ChatGPT-like dialog systems in visual generation.
-
-
-
+
+
+### FLUX Tuners
+
+
+
+ | Yarn Style |
+ Soft Watercolor Style |
+
+
+  |
+  |
+  |
+  |
+  |
+  |
+
+
+ | Travel Style |
+ WuKong Style |
+
+
+  |
+  |
+  |
+  |
+  |
+  |
+
+
+
+
+### ComfyUI Workflow
+
+
+
+
+
+ | Example Workflow Case |
+
+
+
+
+
+
+ |
+
+
+
+
+ |
+
+
+
+
+ |
+
+
+
+
+ |
+
+
+
## 🛠️ Installation
@@ -75,16 +139,17 @@ pip install scepter
### Currently supported approaches
-| Tasks | Methods | Links |
-|:----------------------------:|:--------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
-| Text-to-image generation | SD v1.5 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
-| Text-to-image generation | SD v2.1 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
-| Text-to-image generation | SD-XL | [](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
-| Efficient Tuning | LoRA | [](https://arxiv.org/abs/2106.09685) |
-| Efficient Tuning | Res-Tuning(NeurIPS23) | [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/) |
-| Controllable image synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/) |
-| Image editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [](https://arxiv.org/abs/2403.19534) [](https://ali-vilab.github.io/largen-page/) |
-| Image editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [](https://arxiv.org/abs/2404.12154) [](https://ali-vilab.github.io/stylebooth-page/) |
+| Tasks | Methods | Links |
+|:----------------------------:|:-------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
+| Text-to-image Generation | SD v1.5 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
+| Text-to-image Generation | SD v2.1 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
+| Text-to-image Generation | SD-XL | [](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
+| Efficient Tuning | LoRA | [](https://arxiv.org/abs/2106.09685) |
+| Efficient Tuning | Res-Tuning(NeurIPS23) | [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/) |
+| Controllable Image Synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/) |
+| Image Editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [](https://arxiv.org/abs/2403.19534) [](https://ali-vilab.github.io/largen-page/) |
+| Image Editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [](https://arxiv.org/abs/2404.12154) [](https://ali-vilab.github.io/stylebooth-page/) |
+| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [](https://arxiv.org/abs/2410.00086) [](https://ali-vilab.github.io/ace-page/) |
## 🖥️ SCEPTER Studio
@@ -118,6 +183,22 @@ Therefore, subsequent startups will become much faster (about one minute) as dow
We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/iic/scepter_studio/summary) and [hf_scepter_studio](https://huggingface.co/spaces/modelscope/scepter_studio)
+
+## ⚙️️ ComfyUI Workflow
+
+### Launch
+
+Manually install by moving custom_nodes to ComfyUI.
+```shell
+cd path/to/scepter
+pip install -e .
+cp -r path/to/scepter/workflow/ path/to/ComfyUI/custom_nodes/ComfyUI-Scepter
+cd path/to/ComfyUI
+python main.py
+```
+Alternatively, we will support ComfyUI Manager shortly.
+
+
## 🔍 Learn More
- [Alibaba TongYi Vision Intelligence Lab](https://github.com/ali-vilab)
@@ -150,4 +231,4 @@ This project is licensed under the [Apache License (Version 2.0)](https://github
## Acknowledgement
-Thanks to [Stability-AI](https://github.com/Stability-AI), [SWIFT library](https://github.com/modelscope/swift/) and [Fooocus](https://github.com/lllyasviel/Fooocus) for their awesome work.
+Thanks to [Stability-AI](https://github.com/Stability-AI), [SWIFT library](https://github.com/modelscope/swift/), [Fooocus](https://github.com/lllyasviel/Fooocus) and [ComfyUI](https://github.com/comfyanonymous/ComfyUI) for their awesome work.
\ No newline at end of file
diff --git a/requirements/framework.txt b/requirements/framework.txt
index f01dc07..44c7dd1 100644
--- a/requirements/framework.txt
+++ b/requirements/framework.txt
@@ -2,8 +2,8 @@ albumentations
beautifulsoup4
bezier
einops
-modelscope==1.14.0
-ms-swift>=2.0.1
+modelscope
+ms-swift
numpy
open_clip_torch
opencv-python
@@ -14,3 +14,4 @@ pyyaml>=5.3.1
scikit-image
torchsde
transformers
+scikit-learn
\ No newline at end of file
diff --git a/scepter/methods/examples/generation/dit_flux1.0_dev_1024_lora.yaml b/scepter/methods/examples/generation/dit_flux1.0_dev_1024_lora.yaml
new file mode 100644
index 0000000..fcd03bd
--- /dev/null
+++ b/scepter/methods/examples/generation/dit_flux1.0_dev_1024_lora.yaml
@@ -0,0 +1,303 @@
+ENV:
+ BACKEND: nccl
+ SEED: 166666
+SOLVER:
+ # NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
+ NAME: LatentDiffusionSolver
+ # MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
+ MAX_STEPS: 100000
+ # USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
+ USE_AMP: True
+ # DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
+ DTYPE: bfloat16
+ # USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
+ USE_FAIRSCALE: False
+ # USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
+ USE_FSDP: True
+ # LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
+ LOAD_MODEL_ONLY: False
+ # RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
+ RESUME_FROM:
+ WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
+ LOG_FILE: std_log.txt
+ # EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
+ EVAL_INTERVAL: 100
+ # LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
+ LOG_TRAIN_NUM: 16
+ # FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
+ FSDP_REDUCE_DTYPE: float32
+ # FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
+ FSDP_BUFFER_DTYPE: float32
+ # FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
+ FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
+ SAVE_MODULES: [ 'model'] #
+ TRAIN_MODULES: ['model']
+ #
+ FILE_SYSTEM:
+ NAME: "ModelscopeFs"
+ TEMP_DIR: "./cache/cache_data"
+ #
+ FREEZE:
+ TUNER:
+ - NAME: SwiftLoRA
+ R: 4
+ LORA_ALPHA: 4
+ LORA_DROPOUT: 0.0
+ BIAS: "none"
+ TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
+ ##
+ MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ # NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ # NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
+ NOISE_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
+ NAME: FlowMatchSigmaScheduler
+ # WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ # LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
+ LOGIT_MEAN: 0.0
+ # LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
+ LOGIT_STD: 1.0
+ # MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
+ NAME: FlowMatchFluxShiftScheduler
+ # SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
+ SHIFT: False
+ # SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
+ SIGMOID_SCALE: 1
+ # BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
+ BASE_SHIFT: 0.5
+ # MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'Flux'
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
+ # IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
+ IN_CHANNELS: 64
+ # HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
+ HIDDEN_SIZE: 3072
+ # NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
+ NUM_HEADS: 24
+ # AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
+ AXES_DIM: [ 16, 56, 56 ]
+ # THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
+ THETA: 10000
+ # VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
+ VEC_IN_DIM: 768
+ # GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
+ GUIDANCE_EMBED: True
+ # CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
+ CONTEXT_IN_DIM: 4096
+ # MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
+ MLP_RATIO: 4.0
+ # QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
+ QKV_BIAS: True
+ # DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
+ DEPTH: 19
+ # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
+ DEPTH_SINGLE_BLOCKS: 38
+ USE_GRAD_CHECKPOINT: True
+
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
+ NAME: T5PlusClipFluxEmbedder
+ # T5_MODEL DESCRIPTION: TYPE: default: ''
+ T5_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: T5EncoderModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: T5Tokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 512
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: last_hidden_state
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: False
+ CLEAN: whitespace
+ # CLIP_MODEL DESCRIPTION: TYPE: default: ''
+ CLIP_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: CLIPTextModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 77
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: pooler_output
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: True
+ CLEAN: whitespace
+ USE_GRAD_CHECKPOINT: True
+ #
+ SAMPLE_ARGS:
+ SAMPLE_STEPS: 50
+ SAMPLER: flow_eluer
+ SEED: 2024
+ IMAGE_SIZE: [ 1024, 1024 ]
+ GUIDE_SCALE: 3.5
+ #
+ OPTIMIZER:
+ NAME: AdamW
+ LEARNING_RATE: 4e-4
+ BETAS: [ 0.9, 0.999 ]
+ EPS: 1e-8
+ WEIGHT_DECAY: 1e-2
+ AMSGRAD: False
+ #
+ TRAIN_DATA:
+ NAME: ImageTextPairMSDataset
+ MODE: train
+ MS_DATASET_NAME: style_custom_dataset
+ MS_DATASET_NAMESPACE: damo
+ MS_DATASET_SUBNAME: 3D
+ PROMPT_PREFIX: ""
+ MS_DATASET_SPLIT: train
+ MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
+ REPLACE_STYLE: False
+ PIN_MEMORY: True
+ BATCH_SIZE: 1
+ NUM_WORKERS: 4
+ SAMPLER:
+ NAME: LoopSampler
+ TRANSFORMS:
+ - NAME: LoadImageFromFile
+ RGB_ORDER: RGB
+ BACKEND: pillow
+ - NAME: FlexibleResize
+ INTERPOLATION: bilinear
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: FlexibleCenterCrop
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: ImageToTensor
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: Normalize
+ MEAN: [ 0.5, 0.5, 0.5 ]
+ STD: [ 0.5, 0.5, 0.5 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'image' ]
+ BACKEND: torchvision
+ - NAME: Select
+ KEYS: [ 'image', 'prompt' ]
+ META_KEYS: [ 'data_key' ]
+ #
+ EVAL_DATA:
+ NAME: Text2ImageDataset
+ MODE: eval
+ PROMPT_FILE:
+ PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
+ IMAGE_SIZE: [ 1024, 1024 ]
+ FIELDS: [ "prompt" ]
+ DELIMITER: '#;#'
+ PROMPT_PREFIX: ''
+ PIN_MEMORY: True
+ BATCH_SIZE: 2
+ NUM_WORKERS: 4
+ TRANSFORMS:
+ - NAME: Select
+ KEYS: [ 'index', 'prompt' ]
+ META_KEYS: [ 'image_size' ]
+ #
+ TRAIN_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ PRIORITY: 0
+ - NAME: BackwardHook
+# GRADIENT_CLIP: 1.0
+ PRIORITY: 10
+ - NAME: LogHook
+ LOG_INTERVAL: 10
+ -
+ NAME: TensorboardLogHook
+ -
+ NAME: CheckpointHook
+ INTERVAL: 10000
+ PRIORITY: 200
+ SAVE_LAST: True
+ SAVE_NAME_PREFIX: 'step'
+ DISABLE_SNAPSHOT: True
+ EVAL_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ PRIORITY: 0
\ No newline at end of file
diff --git a/scepter/methods/examples/generation/dit_flux1.0_schnell_1024_lora.yaml b/scepter/methods/examples/generation/dit_flux1.0_schnell_1024_lora.yaml
new file mode 100644
index 0000000..09a4d78
--- /dev/null
+++ b/scepter/methods/examples/generation/dit_flux1.0_schnell_1024_lora.yaml
@@ -0,0 +1,302 @@
+ENV:
+ BACKEND: nccl
+ SEED: 166666
+SOLVER:
+ # NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
+ NAME: LatentDiffusionSolver
+ # MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
+ MAX_STEPS: 100000
+ # USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
+ USE_AMP: True
+ # DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
+ DTYPE: bfloat16
+ # USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
+ USE_FAIRSCALE: False
+ # USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
+ USE_FSDP: True
+ # LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
+ LOAD_MODEL_ONLY: False
+ # RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
+ RESUME_FROM:
+ WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
+ LOG_FILE: std_log.txt
+ # EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
+ EVAL_INTERVAL: 100
+ # LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
+ LOG_TRAIN_NUM: 16
+ # FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
+ FSDP_REDUCE_DTYPE: float32
+ # FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
+ FSDP_BUFFER_DTYPE: float32
+ # FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
+ FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
+ SAVE_MODULES: [ 'model'] #
+ TRAIN_MODULES: ['model']
+ #
+ FILE_SYSTEM:
+ NAME: "ModelscopeFs"
+ TEMP_DIR: "./cache/cache_data"
+ #
+ FREEZE:
+ TUNER:
+ - NAME: SwiftLoRA
+ R: 4
+ LORA_ALPHA: 4
+ LORA_DROPOUT: 0.0
+ BIAS: "none"
+ TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
+ ##
+ MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ # NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ # NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
+ NOISE_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
+ NAME: FlowMatchSigmaScheduler
+ # WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ # LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
+ LOGIT_MEAN: 0.0
+ # LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
+ LOGIT_STD: 1.0
+ # MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
+ NAME: FlowMatchFluxShiftScheduler
+ # SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
+ SHIFT: False
+ # SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
+ SIGMOID_SCALE: 1
+ # BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
+ BASE_SHIFT: 0.5
+ # MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'Flux'
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
+ # IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
+ IN_CHANNELS: 64
+ # HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
+ HIDDEN_SIZE: 3072
+ # NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
+ NUM_HEADS: 24
+ # AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
+ AXES_DIM: [ 16, 56, 56 ]
+ # THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
+ THETA: 10000
+ # VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
+ VEC_IN_DIM: 768
+ # GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
+ GUIDANCE_EMBED: False
+ # CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
+ CONTEXT_IN_DIM: 4096
+ # MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
+ MLP_RATIO: 4.0
+ # QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
+ QKV_BIAS: True
+ # DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
+ DEPTH: 19
+ # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
+ DEPTH_SINGLE_BLOCKS: 38
+ USE_GRAD_CHECKPOINT: True
+
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
+ NAME: T5PlusClipFluxEmbedder
+ # T5_MODEL DESCRIPTION: TYPE: default: ''
+ T5_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: T5EncoderModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: T5Tokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 256
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: last_hidden_state
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: False
+ CLEAN: whitespace
+ # CLIP_MODEL DESCRIPTION: TYPE: default: ''
+ CLIP_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: CLIPTextModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 77
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: pooler_output
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: True
+ CLEAN: whitespace
+ #
+ SAMPLE_ARGS:
+ SAMPLE_STEPS: 4
+ SAMPLER: flow_eluer
+ SEED: 2024
+ IMAGE_SIZE: [ 1024, 1024 ]
+ GUIDE_SCALE: 3.5
+ #
+ OPTIMIZER:
+ NAME: AdamW
+ LEARNING_RATE: 4e-4
+ BETAS: [ 0.9, 0.999 ]
+ EPS: 1e-8
+ WEIGHT_DECAY: 1e-2
+ AMSGRAD: False
+ #
+ TRAIN_DATA:
+ NAME: ImageTextPairMSDataset
+ MODE: train
+ MS_DATASET_NAME: style_custom_dataset
+ MS_DATASET_NAMESPACE: damo
+ MS_DATASET_SUBNAME: 3D
+ PROMPT_PREFIX: ""
+ MS_DATASET_SPLIT: train
+ MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
+ REPLACE_STYLE: False
+ PIN_MEMORY: True
+ BATCH_SIZE: 1
+ NUM_WORKERS: 4
+ SAMPLER:
+ NAME: LoopSampler
+ TRANSFORMS:
+ - NAME: LoadImageFromFile
+ RGB_ORDER: RGB
+ BACKEND: pillow
+ - NAME: FlexibleResize
+ INTERPOLATION: bilinear
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: FlexibleCenterCrop
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: ImageToTensor
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: Normalize
+ MEAN: [ 0.5, 0.5, 0.5 ]
+ STD: [ 0.5, 0.5, 0.5 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'image' ]
+ BACKEND: torchvision
+ - NAME: Select
+ KEYS: [ 'image', 'prompt' ]
+ META_KEYS: [ 'data_key' ]
+ #
+ EVAL_DATA:
+ NAME: Text2ImageDataset
+ MODE: eval
+ PROMPT_FILE:
+ PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
+ IMAGE_SIZE: [ 1024, 1024 ]
+ FIELDS: [ "prompt" ]
+ DELIMITER: '#;#'
+ PROMPT_PREFIX: ''
+ PIN_MEMORY: True
+ BATCH_SIZE: 2
+ NUM_WORKERS: 4
+ TRANSFORMS:
+ - NAME: Select
+ KEYS: [ 'index', 'prompt' ]
+ META_KEYS: [ 'image_size' ]
+ #
+ TRAIN_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ PRIORITY: 0
+ - NAME: BackwardHook
+# GRADIENT_CLIP: 1.0
+ PRIORITY: 10
+ - NAME: LogHook
+ LOG_INTERVAL: 10
+ -
+ NAME: TensorboardLogHook
+ -
+ NAME: CheckpointHook
+ INTERVAL: 10000
+ PRIORITY: 200
+ SAVE_LAST: True
+ SAVE_NAME_PREFIX: 'step'
+ DISABLE_SNAPSHOT: True
+ EVAL_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ PRIORITY: 0
\ No newline at end of file
diff --git a/scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml b/scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml
new file mode 100644
index 0000000..cae7cd1
--- /dev/null
+++ b/scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml
@@ -0,0 +1,201 @@
+NAME: FLUX1.0_DEV
+IS_DEFAULT: False
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[1024, 1024]]
+ INPUT:
+ IMAGE:
+ ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
+ TARGET_SIZE_AS_TUPLE: [1024, 1024]
+ PROMPT: ""
+ NEGATIVE_PROMPT:
+ DEFAULT: ""
+ VISIBLE: False
+ PROMPT_PREFIX: ""
+ SAMPLE:
+ VALUES: ["flow_eluer"]
+ DEFAULT: "flow_eluer"
+ SAMPLE_STEPS: 50
+ GUIDE_SCALE: 3.5
+ GUIDE_RESCALE:
+ DEFAULT: 0.0
+ VISIBLE: False
+ DISCRETIZATION:
+ VALUES: []
+ DEFAULT:
+ VISIBLE: False
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: bfloat16
+ INPUT: ["LATENT"]
+ PARAS:
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: bfloat16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["PROMPT"]
+#
+MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ # NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ # NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
+ NOISE_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
+ NAME: FlowMatchSigmaScheduler
+ # WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ # LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
+ LOGIT_MEAN: 0.0
+ # LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
+ LOGIT_STD: 1.0
+ # MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
+ MODE_SCALE: 1.29
+ #
+ DIFFUSION_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'Flux'
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
+ # IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
+ IN_CHANNELS: 64
+ # HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
+ HIDDEN_SIZE: 3072
+ # NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
+ NUM_HEADS: 24
+ # AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
+ AXES_DIM: [ 16, 56, 56 ]
+ # THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
+ THETA: 10000
+ # VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
+ VEC_IN_DIM: 768
+ # GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
+ GUIDANCE_EMBED: True
+ # CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
+ CONTEXT_IN_DIM: 4096
+ # MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
+ MLP_RATIO: 4.0
+ # QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
+ QKV_BIAS: True
+ # DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
+ DEPTH: 19
+ # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
+ DEPTH_SINGLE_BLOCKS: 38
+
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
+ NAME: T5PlusClipFluxEmbedder
+ # T5_MODEL DESCRIPTION: TYPE: default: ''
+ T5_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: T5EncoderModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: T5Tokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 512
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: last_hidden_state
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: False
+ CLEAN: whitespace
+ # CLIP_MODEL DESCRIPTION: TYPE: default: ''
+ CLIP_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: CLIPTextModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 77
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: pooler_output
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: True
+ CLEAN: whitespace
\ No newline at end of file
diff --git a/scepter/methods/studio/inference/dit/flux1.0_schnell_pro.yaml b/scepter/methods/studio/inference/dit/flux1.0_schnell_pro.yaml
new file mode 100644
index 0000000..c67d9d9
--- /dev/null
+++ b/scepter/methods/studio/inference/dit/flux1.0_schnell_pro.yaml
@@ -0,0 +1,212 @@
+NAME: FLUX1.0_SCHNELL
+IS_DEFAULT: False
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[1024, 1024]]
+ INPUT:
+ IMAGE:
+ ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
+ TARGET_SIZE_AS_TUPLE: [1024, 1024]
+ PROMPT: ""
+ NEGATIVE_PROMPT:
+ DEFAULT: ""
+ VISIBLE: False
+ PROMPT_PREFIX: ""
+ SAMPLE:
+ VALUES: ["flow_eluer"]
+ DEFAULT: "flow_eluer"
+ SAMPLE_STEPS: 4
+ GUIDE_SCALE: 3.5
+ GUIDE_RESCALE:
+ DEFAULT: 0.0
+ VISIBLE: False
+ DISCRETIZATION:
+ VALUES: []
+ DEFAULT:
+ VISIBLE: False
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: bfloat16
+ INPUT: ["LATENT"]
+ PARAS:
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: bfloat16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["PROMPT"]
+#
+MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ # NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ # NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
+ NOISE_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
+ NAME: FlowMatchSigmaScheduler
+ # WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ # LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
+ LOGIT_MEAN: 0.0
+ # LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
+ LOGIT_STD: 1.0
+ # MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
+ NAME: FlowMatchFluxShiftScheduler
+ # SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
+ SHIFT: False
+ # SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
+ SIGMOID_SCALE: 1
+ # BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
+ BASE_SHIFT: 0.5
+ # MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'Flux'
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
+ # IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
+ IN_CHANNELS: 64
+ # HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
+ HIDDEN_SIZE: 3072
+ # NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
+ NUM_HEADS: 24
+ # AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
+ AXES_DIM: [ 16, 56, 56 ]
+ # THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
+ THETA: 10000
+ # VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
+ VEC_IN_DIM: 768
+ # GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
+ GUIDANCE_EMBED: False
+ # CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
+ CONTEXT_IN_DIM: 4096
+ # MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
+ MLP_RATIO: 4.0
+ # QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
+ QKV_BIAS: True
+ # DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
+ DEPTH: 19
+ # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
+ DEPTH_SINGLE_BLOCKS: 38
+
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
+ NAME: T5PlusClipFluxEmbedder
+ # T5_MODEL DESCRIPTION: TYPE: default: ''
+ T5_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: T5EncoderModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: T5Tokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 256
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: last_hidden_state
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: False
+ CLEAN: whitespace
+ # CLIP_MODEL DESCRIPTION: TYPE: default: ''
+ CLIP_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: CLIPTextModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 77
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: pooler_output
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: True
+ CLEAN: whitespace
\ No newline at end of file
diff --git a/scepter/methods/studio/inference/inference.yaml b/scepter/methods/studio/inference/inference.yaml
index e6bf473..5d93331 100644
--- a/scepter/methods/studio/inference/inference.yaml
+++ b/scepter/methods/studio/inference/inference.yaml
@@ -6,65 +6,81 @@ DIFFUSION_PARAS:
'dpm2_karras', 'dpm2_ancestral_karras', 'dpmpp_2s_ancestral_karras', 'dpmpp_2m_karras',
'dpmpp_sde_karras', 'dpmpp_2m_sde_karras']
DEFAULT: 'dpmpp_2s_ancestral'
+ VISIBLE: True
NEGATIVE_PROMPT:
DEFAULT:
+ VISIBLE: True
PROMPT_PREFIX:
DEFAULT:
+ VISIBLE: True
SAMPLES:
MIN: 1
MAX: 4
DEFAULT: 1
+ VISIBLE: True
SAMPLE_STEPS:
MIN: 1
MAX: 100
DEFAULT: 30
+ VISIBLE: True
GUIDE_SCALE:
MIN: 0
MAX: 10
DEFAULT: 5.0
+ VISIBLE: True
GUIDE_RESCALE:
MIN: 0
MAX: 1.0
DEFAULT: 0.5
+ VISIBLE: True
DISCRETIZATION:
VALUES: ["trailing", "leading", "linspace"]
DEFAULT: "linspace"
+ VISIBLE: True
REFINE_SAMPLERS:
VALUES: [ 'ddim', 'euler', 'euler_ancestral', 'heun', 'dpm2',
'dpm2_ancestral', 'dpmpp_2m', 'dpmpp_sde', 'dpmpp_2m_sde', 'dpmpp_2s_ancestral',
'dpm2_karras', 'dpm2_ancestral_karras', 'dpmpp_2s_ancestral_karras', 'dpmpp_2m_karras',
'dpmpp_sde_karras', 'dpmpp_2m_sde_karras' ]
DEFAULT: 'dpmpp_2s_ancestral'
+ VISIBLE: True
REFINE_SAMPLE_STEPS:
MIN: 0
MAX: 100
DEFAULT: 30
+ VISIBLE: True
REFINE_GUIDE_SCALE:
MIN: 0
MAX: 10
DEFAULT: 5.0
+ VISIBLE: True
REFINE_GUIDE_RESCALE:
MIN: 0
MAX: 1.0
DEFAULT: 0.5
+ VISIBLE: True
REFINE_DISCRETIZATION:
VALUES: [ "trailing", "leading", "linspace" ]
DEFAULT: "linspace"
+ VISIBLE: True
AESTHETIC_SCORE:
MIN: 0.0
MAX: 10.0
DEFAULT: 6.0
+ VISIBLE: True
NEGATIVE_AESTHETIC_SCORE:
MIN: 0.0
MAX: 10.0
DEFAULT: 2.5
+ VISIBLE: True
REFINE_STRENGTH:
MIN: 0
MAX: 1.0
DEFAULT: 0.15
RESOLUTIONS:
- VALUES: [ [512, 512], [768, 768],
- [704, 1408], [704, 1344], [768, 1344],
+ VALUES: [
+ [512, 512], [768, 768],
+ [704, 1408], [704, 1344], [768, 1344],
[720, 1280],
[768, 1280], [832, 1216], [832, 1152],
[896, 1152], [896, 1088], [960, 1088],
@@ -74,8 +90,13 @@ DIFFUSION_PARAS:
[1280, 768],
[1344, 768], [1344, 704], [1408, 704],
[1472, 704], [1536, 640], [1600, 640],
- [1664, 576], [1728, 576]]
+ [1664, 576], [1728, 576],
+ [2048, 2048], [2048, 1920], [1920, 2048],
+ [1536, 2560], [2560, 1536], [2560, 1440],
+ [2560, 1440]
+ ]
DEFAULT: [1024, 1024]
+ VISIBLE: True
EXTENSION_PARAS:
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
diff --git a/scepter/methods/studio/preprocess/preprocess.yaml b/scepter/methods/studio/preprocess/preprocess.yaml
index 8e43da0..68eabe0 100644
--- a/scepter/methods/studio/preprocess/preprocess.yaml
+++ b/scepter/methods/studio/preprocess/preprocess.yaml
@@ -2,11 +2,10 @@ WORK_DIR: datasets
EXPORT_DIR: export_datasets
FILE_SYSTEM:
-
- # NAME DESCRIPTION: TYPE: default: ''
NAME: LocalFs
AUTO_CLEAN: False
-
PROCESSORS:
+ # Caption processor
- NAME: BlipImageBase
TYPE: caption
MODEL_PATH: ms://cubeai/blip-image-captioning-base
@@ -15,6 +14,81 @@ PROCESSORS:
PARAS:
- LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
+ - NAME: InternVL15
+ TYPE: caption
+ MODEL_PATH: ms://AI-ModelScope/InternVL-Chat-V1-5
+ DEVICE: "gpu"
+ MEMORY: 49968
+ PARAS:
+ - PROMPT: 用中文描述这张图片
+ LANGUAGE_NAME: Chinese
+ LANGUAGE_ZH_NAME: 中文
+ - PROMPT: Generate the caption in English
+ LANGUAGE_NAME: English
+ LANGUAGE_ZH_NAME: 英语
+ - NAME: QWVLQuantize
+ TYPE: caption
+ DEVICE: "gpu"
+ MEMORY: 7885
+ MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
+ PARAS:
+ - PROMPT: 用中文描述这张图片
+ LANGUAGE_NAME: Chinese
+ LANGUAGE_ZH_NAME: 中文
+ MAX_NEW_TOKENS:
+ VALUE: 1024
+ MAX: 2048
+ STEP: 128
+ MIN: 256
+ MIN_NEW_TOKENS:
+ VALUE: 16
+ MAX: 1024
+ STEP: 16
+ MIN: 0
+ NUM_BEAMS:
+ VALUE: 1
+ MAX: 12
+ STEP: 1
+ MIN: 1
+ REPETITION_PENALTY:
+ VALUE: 1.0
+ MAX: 100.0
+ STEP: 1.0
+ MIN: 1.0
+ TEMPERATURE:
+ VALUE: 1.0
+ MAX: 100.0
+ STEP: 1.0
+ MIN: 1.0
+ - PROMPT: Generate the caption in English
+ LANGUAGE_NAME: English
+ LANGUAGE_ZH_NAME: 英语
+ MAX_NEW_TOKENS:
+ VALUE: 1024
+ MAX: 2048
+ STEP: 128
+ MIN: 256
+ MIN_NEW_TOKENS:
+ VALUE: 16
+ MAX: 1024
+ STEP: 16
+ MIN: 0
+ NUM_BEAMS:
+ VALUE: 1
+ MAX: 12
+ STEP: 1
+ MIN: 1
+ REPETITION_PENALTY:
+ VALUE: 1.0
+ MAX: 100.0
+ STEP: 1.0
+ MIN: 1.0
+ TEMPERATURE:
+ VALUE: 1.0
+ MAX: 100.0
+ STEP: 1.0
+ MIN: 1.0
+
- NAME: QWVL
TYPE: caption
MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
@@ -77,75 +151,14 @@ PROCESSORS:
MAX: 100.0
STEP: 1.0
MIN: 1.0
- -
- NAME: QWVLQuantize
- TYPE: caption
- DEVICE: "gpu"
- MEMORY: 7885
- MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
- PARAS:
- - PROMPT: 用中文描述这张图片
- LANGUAGE_NAME: Chinese
- LANGUAGE_ZH_NAME: 中文
- MAX_NEW_TOKENS:
- VALUE: 1024
- MAX: 2048
- STEP: 128
- MIN: 256
- MIN_NEW_TOKENS:
- VALUE: 16
- MAX: 1024
- STEP: 16
- MIN: 0
- NUM_BEAMS:
- VALUE: 1
- MAX: 12
- STEP: 1
- MIN: 1
- REPETITION_PENALTY:
- VALUE: 1.0
- MAX: 100.0
- STEP: 1.0
- MIN: 1.0
- TEMPERATURE:
- VALUE: 1.0
- MAX: 100.0
- STEP: 1.0
- MIN: 1.0
- - PROMPT: Generate the caption in English
- LANGUAGE_NAME: English
- LANGUAGE_ZH_NAME: 英语
- MAX_NEW_TOKENS:
- VALUE: 1024
- MAX: 2048
- STEP: 128
- MIN: 256
- MIN_NEW_TOKENS:
- VALUE: 16
- MAX: 1024
- STEP: 16
- MIN: 0
- NUM_BEAMS:
- VALUE: 1
- MAX: 12
- STEP: 1
- MIN: 1
- REPETITION_PENALTY:
- VALUE: 1.0
- MAX: 100.0
- STEP: 1.0
- MIN: 1.0
- TEMPERATURE:
- VALUE: 1.0
- MAX: 100.0
- STEP: 1.0
- MIN: 1.0
- -
- NAME: CenterCrop
+
+ # Simple processor
+ - NAME: CenterCrop
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
+ CAPTION_INTERACTIVE: False
HEIGHT_RATIO:
VALUE: 1
MAX: 20
@@ -156,18 +169,284 @@ PROCESSORS:
MAX: 20
STEP: 1
MIN: 1
-# - NAME: PaddingCrop
-# TYPE: image
-# DEVICE: "cpu"
-# MEMORY: 10
-# PARAS:
-# HEIGHT_RATIO:
-# VALUE: 3
-# MAX: 25
-# STEP: 1
-# MIN: 1
-# WIDTH_RATIO:
-# VALUE: 4
-# MAX: 20
-# STEP: 1
-# MIN: 1
+ - NAME: ChangeSample
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ CAPTION_INTERACTIVE: False
+ SRC_IMAGE_INTERACTIVE: True
+ SRC_IMAGE_MASK_INTERACTIVE: True
+ TARGET_IMAGE_INTERACTIVE: True
+ PREVIEW_BTN_VISIBLE: False
+ - NAME: MaskEditSample
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ SRC_IMAGE_INTERACTIVE: True
+ SRC_IMAGE_TOOL: sketch
+ CAPTION_INTERACTIVE: False
+ - NAME: SourceMaskSample
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ SRC_IMAGE_INTERACTIVE: True
+ SRC_IMAGE_TOOL: sketch
+ CAPTION_INTERACTIVE: False
+ - NAME: MaskSwapEditSample
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ PREVIEW_BTN_VISIBLE: False
+ SRC_IMAGE_INTERACTIVE: True
+ SRC_IMAGE_TOOL: sketch
+ CAPTION_INTERACTIVE: False
+ - NAME: SwapMaskSwapEditSample
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ PREVIEW_BTN_VISIBLE: False
+ SRC_IMAGE_INTERACTIVE: True
+ SRC_IMAGE_TOOL: sketch
+ CAPTION_INTERACTIVE: False
+ - NAME: SwapSample
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ PREVIEW_BTN_VISIBLE: True
+ SRC_IMAGE_INTERACTIVE: True
+ TARGET_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: DefaultMaskSample
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ PREVIEW_BTN_VISIBLE: True
+ SRC_IMAGE_INTERACTIVE: False
+ TARGET_IMAGE_INTERACTIVE: False
+ CAPTION_INTERACTIVE: False
+ # - NAME: PaddingCrop
+ # TYPE: image
+ # DEVICE: "cpu"
+ # MEMORY: 10
+ # PARAS:
+ # HEIGHT_RATIO:
+ # VALUE: 3
+ # MAX: 25
+ # STEP: 1
+ # MIN: 1
+ # WIDTH_RATIO:
+ # VALUE: 4
+ # MAX: 20
+ # STEP: 1
+ # MIN: 1
+ # Annotator processor
+ - NAME: CannyExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "CannyAnnotator"
+ LOW_THRESHOLD: 100
+ HIGH_THRESHOLD: 200
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: ColorExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "ColorAnnotator"
+ RATIO: 64
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: InfoDrawContourExtractor
+ TYPE: image
+ DEVICE: "gpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "InfoDrawContourAnnotator"
+ INPUT_NC: 3
+ OUTPUT_NC: 1
+ N_RESIDUAL_BLOCKS: 3
+ SIGMOID: True
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_contour_style.pth"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: DegradationExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "DegradationAnnotator"
+ RANDOM_DEGRADATION: True
+ PARAMS:
+ gaussian_noise: { }
+ resize: { 'scale': [ 0.4, 0.8 ] }
+ jpeg: { 'jpeg_level': [ 25, 75 ] }
+ gaussian_blur: { 'kernel_size': [ 7, 9, 11, 13, 15 ], 'sigma': [ 0.9, 1.8 ] }
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: MidasExtractor
+ TYPE: image
+ DEVICE: "gpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "MidasDetector"
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: DoodleExtractor
+ TYPE: image
+ DEVICE: "gpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "DoodleAnnotator"
+ PROCESSOR_TYPE: "pidinet_sketch"
+ PROCESSOR_CFG:
+ - NAME: "PiDiAnnotator"
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/table5_pidinet.pth"
+ - NAME: "SketchAnnotator"
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/sketch_simplification_gan.pth"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: GrayExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "GrayAnnotator"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: InpaintingExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "InpaintingAnnotator"
+ RETURN_MASK: False
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+
+ - NAME: InpaintingSourceExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: OpenposeExtractor
+ TYPE: image
+ DEVICE: "gpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "OpenposeAnnotator"
+ BODY_MODEL_PATH: "ms://iic/scepter_annotator@annotator/ckpts/body_pose_model.pth"
+ HAND_MODEL_PATH: "ms://iic/scepter_annotator@annotator/ckpts/hand_pose_model.pth"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: OutpaintingExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "OutpaintingAnnotator"
+ RETURN_MASK: False
+ KEEP_PADDING_RATIO: 1
+ RANDOM_CFG:
+ DIRECTION_RANGE: [ 'left', 'right', 'up', 'down' ]
+ RATIO_RANGE: [ 0.1, 0.7 ]
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ USE_MASK_VISIBLE: True
+ PREVIEW_BTN_VISIBLE: False
+ - NAME: OutpaintingResize
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "OutpaintingResize"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ PREVIEW_BTN_VISIBLE: False
+ - NAME: InfoDrawAnimeAnnotator
+ TYPE: image
+ DEVICE: "gpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "InfoDrawAnimeAnnotator"
+ INPUT_NC: 3
+ OUTPUT_NC: 1
+ N_RESIDUAL_BLOCKS: 3
+ SIGMOID: True
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_anime_style.pth"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: ESAMExtractor
+ TYPE: image
+ DEVICE: "gpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "ESAMAnnotator"
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/efficient_sam_vits.pt"
+ SAVE_MODE: 'P'
+ GRID_SIZE: 32
+ USE_DOMINANT_COLOR: True
+ RETURN_MASK: False
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+ - NAME: InvertExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "InvertAnnotator"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
+
+ - NAME: LamaExtractor
+ TYPE: image
+ DEVICE: "cpu"
+ MEMORY: 10
+ MODEL:
+ NAME: "LamaAnnotator"
+ PRETRAINED_MODEL: "ms:///iic/cv_fft_inpainting_lama/"
+ PARAS:
+ SRC_IMAGE_TOOL: sketch
+ SRC_IMAGE_INTERACTIVE: True
+ CAPTION_INTERACTIVE: False
diff --git a/scepter/methods/studio/self_train/dit/flux1.0_dv_pro.yaml b/scepter/methods/studio/self_train/dit/flux1.0_dv_pro.yaml
new file mode 100644
index 0000000..23701d1
--- /dev/null
+++ b/scepter/methods/studio/self_train/dit/flux1.0_dv_pro.yaml
@@ -0,0 +1,343 @@
+ENV:
+ BACKEND: nccl
+META:
+ VERSION: 'FLUX1.0_DEV'
+ DESCRIPTION: "flux 1.0 dev"
+ IS_DEFAULT: False
+ IS_SHARE: True
+ INFERENCE_PARAS:
+ INFERENCE_BATCH_SIZE: 1
+ INFERENCE_PREFIX: ""
+ DEFAULT_SAMPLER: "flow_euler"
+ DEFAULT_SAMPLE_STEPS: 50
+ INFERENCE_N_PROMPT: ""
+ RESOLUTION: [1024, 1024]
+ PARAS:
+ -
+ TRAIN_BATCH_SIZE: 1
+ TRAIN_PREFIX: ""
+ TRAIN_N_PROMPT: ""
+ RESOLUTION: [1024, 1024]
+ MEMORY: 89000
+ EPOCHS: 50
+ SAVE_INTERVAL: 25
+ EPSEC: 0.818
+ LEARNING_RATE: 4e-4
+ IS_DEFAULT: False
+ TUNER: FULL
+ -
+ TRAIN_BATCH_SIZE: 1
+ TRAIN_PREFIX: ""
+ TRAIN_N_PROMPT: ""
+ RESOLUTION: [1024, 1024]
+ MEMORY: 89000
+ EPOCHS: 50
+ SAVE_INTERVAL: 25
+ EPSEC: 0.818
+ LEARNING_RATE: 4e-4
+ IS_DEFAULT: True
+ TUNER: LORA
+ #
+ TUNERS:
+ LORA:
+ -
+ NAME: SwiftLoRA
+ R: 4
+ LORA_ALPHA: 4
+ LORA_DROPOUT: 0.0
+ BIAS: "none"
+ TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
+#
+SOLVER:
+ NAME: LatentDiffusionSolver
+ # MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
+ MAX_STEPS: 100000
+ # USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
+ USE_AMP: True
+ # DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
+ DTYPE: bfloat16
+ # USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
+ USE_FAIRSCALE: False
+ # USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
+ USE_FSDP: True
+ # LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
+ LOAD_MODEL_ONLY: False
+ # RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
+ RESUME_FROM:
+ WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
+ LOG_FILE: std_log.txt
+ # EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
+ EVAL_INTERVAL: 100
+ # LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
+ LOG_TRAIN_NUM: 16
+ # FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
+ FSDP_REDUCE_DTYPE: float32
+ # FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
+ FSDP_BUFFER_DTYPE: float32
+ # FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
+ FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
+ SAVE_MODULES: [ 'model'] #
+ TRAIN_MODULES: ['model']
+ #
+ FILE_SYSTEM:
+ NAME: "ModelscopeFs"
+ TEMP_DIR: "./cache/cache_data"
+ #
+ FREEZE:
+ #
+ TUNER:
+ #
+ MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ # NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ # NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
+ NOISE_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
+ NAME: FlowMatchSigmaScheduler
+ # WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ # LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
+ LOGIT_MEAN: 0.0
+ # LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
+ LOGIT_STD: 1.0
+ # MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
+ NAME: FlowMatchFluxShiftScheduler
+ # SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
+ SHIFT: False
+ # SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
+ SIGMOID_SCALE: 1
+ # BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
+ BASE_SHIFT: 0.5
+ # MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'Flux'
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
+ # IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
+ IN_CHANNELS: 64
+ # HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
+ HIDDEN_SIZE: 3072
+ # NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
+ NUM_HEADS: 24
+ # AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
+ AXES_DIM: [ 16, 56, 56 ]
+ # THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
+ THETA: 10000
+ # VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
+ VEC_IN_DIM: 768
+ # GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
+ GUIDANCE_EMBED: False
+ # CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
+ CONTEXT_IN_DIM: 4096
+ # MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
+ MLP_RATIO: 4.0
+ # QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
+ QKV_BIAS: True
+ # DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
+ DEPTH: 19
+ # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
+ DEPTH_SINGLE_BLOCKS: 38
+ USE_GRAD_CHECKPOINT: True
+
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
+ NAME: T5PlusClipFluxEmbedder
+ # T5_MODEL DESCRIPTION: TYPE: default: ''
+ T5_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: T5EncoderModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: T5Tokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 512
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: last_hidden_state
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: False
+ CLEAN: whitespace
+ # CLIP_MODEL DESCRIPTION: TYPE: default: ''
+ CLIP_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: CLIPTextModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 77
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: pooler_output
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: True
+ CLEAN: whitespace
+ #
+ SAMPLE_ARGS:
+ SAMPLE_STEPS: 50
+ SAMPLER: flow_eluer
+ SEED: 2024
+ IMAGE_SIZE: [ 1024, 1024 ]
+ SHIFT: True
+ GUIDE_SCALE: 3.5
+ #
+ OPTIMIZER:
+ NAME: AdamW
+ LEARNING_RATE: 4e-4
+ BETAS: [ 0.9, 0.999 ]
+ EPS: 1e-8
+ WEIGHT_DECAY: 1e-2
+ AMSGRAD: False
+ #
+ TRAIN_DATA:
+ NAME: ImageTextPairMSDataset
+ MODE: train
+ MS_DATASET_NAME: style_custom_dataset
+ MS_DATASET_NAMESPACE: damo
+ MS_DATASET_SUBNAME: 3D
+ PROMPT_PREFIX: ""
+ MS_DATASET_SPLIT: train
+ MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
+ REPLACE_STYLE: False
+ PIN_MEMORY: True
+ BATCH_SIZE: 1
+ NUM_WORKERS: 4
+ SAMPLER:
+ NAME: LoopSampler
+ TRANSFORMS:
+ - NAME: LoadImageFromFile
+ RGB_ORDER: RGB
+ BACKEND: pillow
+ - NAME: FlexibleResize
+ INTERPOLATION: bilinear
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: FlexibleCenterCrop
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: ImageToTensor
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: Normalize
+ MEAN: [ 0.5, 0.5, 0.5 ]
+ STD: [ 0.5, 0.5, 0.5 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'image' ]
+ BACKEND: torchvision
+ - NAME: Select
+ KEYS: [ 'image', 'prompt' ]
+ META_KEYS: [ 'data_key' ]
+ #
+ EVAL_DATA:
+ NAME: Text2ImageDataset
+ MODE: eval
+ PROMPT_FILE:
+ PROMPT_DATA: [ "a cat holds a blackboard that writes \"hello world\"", "a dog running on the lawn" ]
+ IMAGE_SIZE: [ 1024, 1024 ]
+ FIELDS: [ "prompt" ]
+ DELIMITER: '#;#'
+ PROMPT_PREFIX: ''
+ PIN_MEMORY: True
+ BATCH_SIZE: 2
+ NUM_WORKERS: 4
+ TRANSFORMS:
+ - NAME: Select
+ KEYS: [ 'index', 'prompt' ]
+ META_KEYS: [ 'image_size' ]
+ #
+ TRAIN_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ PRIORITY: 0
+ - NAME: BackwardHook
+# GRADIENT_CLIP: 1.0
+ PRIORITY: 10
+ - NAME: LogHook
+ LOG_INTERVAL: 10
+ - NAME: CheckpointHook
+ INTERVAL: 10000
+ PRIORITY: 200
+ SAVE_LAST: True
+ SAVE_NAME_PREFIX: 'step'
+ DISABLE_SNAPSHOT: True
+ EVAL_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ SAVE_LAST: True
+ SAVE_NAME_PREFIX: 'step'
+ SAVE_PROBE_PREFIX: 'image'
diff --git a/scepter/methods/studio/self_train/dit/flux1.0_schnell_pro.yaml b/scepter/methods/studio/self_train/dit/flux1.0_schnell_pro.yaml
new file mode 100644
index 0000000..265285c
--- /dev/null
+++ b/scepter/methods/studio/self_train/dit/flux1.0_schnell_pro.yaml
@@ -0,0 +1,342 @@
+ENV:
+ BACKEND: nccl
+META:
+ VERSION: 'FLUX1.0_SCHNELL'
+ DESCRIPTION: "flux 1.0 schnell"
+ IS_DEFAULT: False
+ IS_SHARE: True
+ INFERENCE_PARAS:
+ INFERENCE_BATCH_SIZE: 1
+ INFERENCE_PREFIX: ""
+ DEFAULT_SAMPLER: "flow_euler"
+ DEFAULT_SAMPLE_STEPS: 4
+ INFERENCE_N_PROMPT: ""
+ RESOLUTION: [1024, 1024]
+ PARAS:
+ -
+ TRAIN_BATCH_SIZE: 1
+ TRAIN_PREFIX: ""
+ TRAIN_N_PROMPT: ""
+ RESOLUTION: [1024, 1024]
+ MEMORY: 89000
+ EPOCHS: 50
+ SAVE_INTERVAL: 25
+ EPSEC: 0.818
+ LEARNING_RATE: 4e-4
+ IS_DEFAULT: False
+ TUNER: FULL
+ -
+ TRAIN_BATCH_SIZE: 1
+ TRAIN_PREFIX: ""
+ TRAIN_N_PROMPT: ""
+ RESOLUTION: [1024, 1024]
+ MEMORY: 89000
+ EPOCHS: 50
+ SAVE_INTERVAL: 25
+ EPSEC: 0.818
+ LEARNING_RATE: 4e-4
+ IS_DEFAULT: True
+ TUNER: LORA
+ #
+ TUNERS:
+ LORA:
+ -
+ NAME: SwiftLoRA
+ R: 4
+ LORA_ALPHA: 4
+ LORA_DROPOUT: 0.0
+ BIAS: "none"
+ TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
+#
+SOLVER:
+ NAME: LatentDiffusionSolver
+ # MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
+ MAX_STEPS: 100000
+ # USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
+ USE_AMP: True
+ # DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
+ DTYPE: bfloat16
+ # USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
+ USE_FAIRSCALE: False
+ # USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
+ USE_FSDP: True
+ # LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
+ LOAD_MODEL_ONLY: False
+ # RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
+ RESUME_FROM:
+ WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
+ LOG_FILE: std_log.txt
+ # EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
+ EVAL_INTERVAL: 100
+ # LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
+ LOG_TRAIN_NUM: 16
+ # FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
+ FSDP_REDUCE_DTYPE: float32
+ # FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
+ FSDP_BUFFER_DTYPE: float32
+ # FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
+ FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
+ SAVE_MODULES: [ 'model'] #
+ TRAIN_MODULES: ['model']
+ #
+ FILE_SYSTEM:
+ NAME: "ModelscopeFs"
+ TEMP_DIR: "./cache/cache_data"
+ #
+ FREEZE:
+ #
+ TUNER:
+ #
+ MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ # NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ # NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
+ NOISE_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
+ NAME: FlowMatchSigmaScheduler
+ # WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ # LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
+ LOGIT_MEAN: 0.0
+ # LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
+ LOGIT_STD: 1.0
+ # MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ # NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
+ NAME: FlowMatchFluxShiftScheduler
+ # SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
+ SHIFT: False
+ # SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
+ SIGMOID_SCALE: 1
+ # BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
+ BASE_SHIFT: 0.5
+ # MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'Flux'
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
+ # IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
+ IN_CHANNELS: 64
+ # HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
+ HIDDEN_SIZE: 3072
+ # NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
+ NUM_HEADS: 24
+ # AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
+ AXES_DIM: [ 16, 56, 56 ]
+ # THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
+ THETA: 10000
+ # VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
+ VEC_IN_DIM: 768
+ # GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
+ GUIDANCE_EMBED: False
+ # CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
+ CONTEXT_IN_DIM: 4096
+ # MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
+ MLP_RATIO: 4.0
+ # QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
+ QKV_BIAS: True
+ # DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
+ DEPTH: 19
+ # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
+ DEPTH_SINGLE_BLOCKS: 38
+ USE_GRAD_CHECKPOINT: True
+
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
+ NAME: T5PlusClipFluxEmbedder
+ # T5_MODEL DESCRIPTION: TYPE: default: ''
+ T5_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: T5EncoderModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: T5Tokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 256
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: last_hidden_state
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: False
+ CLEAN: whitespace
+ # CLIP_MODEL DESCRIPTION: TYPE: default: ''
+ CLIP_MODEL:
+ # NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
+ NAME: HFEmbedder
+ # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_MODEL_CLS: CLIPTextModel
+ # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
+ # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
+ # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
+ MAX_LENGTH: 77
+ # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
+ OUTPUT_KEY: pooler_output
+ # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
+ D_TYPE: bfloat16
+ # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
+ BATCH_INFER: True
+ CLEAN: whitespace
+ #
+ SAMPLE_ARGS:
+ SAMPLE_STEPS: 4
+ SAMPLER: flow_eluer
+ SEED: 2024
+ IMAGE_SIZE: [ 1024, 1024 ]
+ GUIDE_SCALE: 3.5
+ #
+ OPTIMIZER:
+ NAME: AdamW
+ LEARNING_RATE: 4e-4
+ BETAS: [ 0.9, 0.999 ]
+ EPS: 1e-8
+ WEIGHT_DECAY: 1e-2
+ AMSGRAD: False
+ #
+ TRAIN_DATA:
+ NAME: ImageTextPairMSDataset
+ MODE: train
+ MS_DATASET_NAME: style_custom_dataset
+ MS_DATASET_NAMESPACE: damo
+ MS_DATASET_SUBNAME: 3D
+ PROMPT_PREFIX: ""
+ MS_DATASET_SPLIT: train
+ MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
+ REPLACE_STYLE: False
+ PIN_MEMORY: True
+ BATCH_SIZE: 1
+ NUM_WORKERS: 4
+ SAMPLER:
+ NAME: LoopSampler
+ TRANSFORMS:
+ - NAME: LoadImageFromFile
+ RGB_ORDER: RGB
+ BACKEND: pillow
+ - NAME: FlexibleResize
+ INTERPOLATION: bilinear
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: FlexibleCenterCrop
+ SIZE: [ 1024, 1024 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: ImageToTensor
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'img' ]
+ BACKEND: pillow
+ - NAME: Normalize
+ MEAN: [ 0.5, 0.5, 0.5 ]
+ STD: [ 0.5, 0.5, 0.5 ]
+ INPUT_KEY: [ 'img' ]
+ OUTPUT_KEY: [ 'image' ]
+ BACKEND: torchvision
+ - NAME: Select
+ KEYS: [ 'image', 'prompt' ]
+ META_KEYS: [ 'data_key' ]
+ #
+ EVAL_DATA:
+ NAME: Text2ImageDataset
+ MODE: eval
+ PROMPT_FILE:
+ PROMPT_DATA: [ "a cat holds a blackboard that writes \"hello world\"", "a dog running on the lawn" ]
+ IMAGE_SIZE: [ 1024, 1024 ]
+ FIELDS: [ "prompt" ]
+ DELIMITER: '#;#'
+ PROMPT_PREFIX: ''
+ PIN_MEMORY: True
+ BATCH_SIZE: 2
+ NUM_WORKERS: 4
+ TRANSFORMS:
+ - NAME: Select
+ KEYS: [ 'index', 'prompt' ]
+ META_KEYS: [ 'image_size' ]
+ #
+ TRAIN_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ PRIORITY: 0
+ - NAME: BackwardHook
+# GRADIENT_CLIP: 1.0
+ PRIORITY: 10
+ - NAME: LogHook
+ LOG_INTERVAL: 10
+ - NAME: CheckpointHook
+ INTERVAL: 10000
+ PRIORITY: 200
+ SAVE_LAST: True
+ SAVE_NAME_PREFIX: 'step'
+ DISABLE_SNAPSHOT: True
+ EVAL_HOOKS:
+ - NAME: ProbeDataHook
+ PROB_INTERVAL: 100
+ SAVE_LAST: True
+ SAVE_NAME_PREFIX: 'step'
+ SAVE_PROBE_PREFIX: 'image'
diff --git a/scepter/methods/studio/tuner_manager/tuner_manager.yaml b/scepter/methods/studio/tuner_manager/tuner_manager.yaml
index ec29345..32e766a 100644
--- a/scepter/methods/studio/tuner_manager/tuner_manager.yaml
+++ b/scepter/methods/studio/tuner_manager/tuner_manager.yaml
@@ -13,3 +13,5 @@ BASE_MODEL_VERSION:
TUNER_TYPE: [ 'LORA', 'SCE', 'FULL' ]
- BASE_MODEL: 'EDIT'
TUNER_TYPE: [ 'LORA', 'SCE', 'FULL' ]
+ - BASE_MODEL: 'FLUX1.0_DEV'
+ TUNER_TYPE: [ 'LORA' ]
\ No newline at end of file
diff --git a/scepter/modules/annotator/__init__.py b/scepter/modules/annotator/__init__.py
index f14a40f..7232daf 100644
--- a/scepter/modules/annotator/__init__.py
+++ b/scepter/modules/annotator/__init__.py
@@ -3,9 +3,21 @@
from scepter.modules.annotator.base_annotator import GeneralAnnotator
from scepter.modules.annotator.canny import CannyAnnotator
from scepter.modules.annotator.color import ColorAnnotator
+from scepter.modules.annotator.degradation import DegradationAnnotator
+from scepter.modules.annotator.doodle import DoodleAnnotator
+from scepter.modules.annotator.gray import GrayAnnotator
from scepter.modules.annotator.hed import HedAnnotator
from scepter.modules.annotator.identity import IdentityAnnotator
+from scepter.modules.annotator.informative_drawing import (
+ InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
+ InfoDrawOpenSketchAnnotator)
+from scepter.modules.annotator.inpainting import InpaintingAnnotator
from scepter.modules.annotator.invert import InvertAnnotator
from scepter.modules.annotator.midas_op import MidasDetector
from scepter.modules.annotator.mlsd_op import MLSDdetector
from scepter.modules.annotator.openpose import OpenposeAnnotator
+from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
+from scepter.modules.annotator.pidinet import PiDiAnnotator
+from scepter.modules.annotator.segmentation import ESAMAnnotator
+from scepter.modules.annotator.sketch import SketchAnnotator
+from scepter.modules.annotator.lama import LamaAnnotator
diff --git a/scepter/modules/annotator/degradation.py b/scepter/modules/annotator/degradation.py
new file mode 100644
index 0000000..7b1ae08
--- /dev/null
+++ b/scepter/modules/annotator/degradation.py
@@ -0,0 +1,135 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import math
+import random
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+
+from PIL import Image
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import Config, dict_to_yaml
+
+
+def gaussian_noise_op(im, v):
+ from basicsr.data.degradations import random_add_gaussian_noise
+ noise_level = v.get('noise_level', [10, 20])
+ out = random_add_gaussian_noise(
+ im,
+ sigma_range=noise_level,
+ clip=True,
+ rounds=False,
+ gray_prob=0.4,
+ )
+ out = np.clip(out, 0.0, 1.0)
+ return out
+
+
+def resize_op(im, v):
+ scale = v.get('scale', [0.5, 0.8])
+ h, w = im.shape[:2]
+ scale = random.uniform(scale[0], scale[1])
+ h_, w_ = int(h * scale), int(w * scale)
+ mode = v.get('mode', 'nearest')
+ if mode == 'nearest':
+ interpolation = cv2.INTER_NEAREST
+ elif mode == 'bilinear':
+ interpolation = cv2.INTER_LINEAR
+ elif mode == 'bicubic':
+ interpolation = cv2.INTER_CUBIC
+ else:
+ interpolation = cv2.INTER_NEAREST
+ im = cv2.resize(im, (w_, h_), interpolation=interpolation)
+ out = cv2.resize(im, (w, h), interpolation=interpolation)
+ out = np.clip(out, 0.0, 1.0)
+ return out
+
+
+def jpeg_op(im, v):
+ from basicsr.data.degradations import add_jpg_compression
+ jpeg_level = v.get('jpeg_level', [50, 75])
+ v = int(random.uniform(jpeg_level[0], jpeg_level[1]))
+ out = add_jpg_compression(im, v)
+ out = np.clip(out, 0.0, 1.0)
+ return out
+
+
+def gaussian_blur_op(im, v):
+ from basicsr.data.degradations import random_mixed_kernels
+ kernel_range = v.get('kernel_size', [7, 9])
+ kernel_size = random.choice(kernel_range)
+ kernel_size = min(int(kernel_size) // 2 * 2 + 1, 21)
+ blur_sigma = v.get('sigma', [0.9, 1.0])
+ kernel = random_mixed_kernels(
+ ('iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso',
+ 'plateau_aniso'), (0.45, 0.25, 0.12, 0.03, 0.12, 0.03),
+ kernel_size,
+ blur_sigma,
+ blur_sigma, [-math.pi, math.pi], [0.5, 2.0], [1, 1.5],
+ noise_range=None)
+
+ pad_size = (21 - kernel_size) // 2
+ kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size)))
+ out = cv2.filter2D(im, -1, kernel)
+ out = np.clip(out, 0.0, 1.0)
+ return out
+
+
+@ANNOTATORS.register_class()
+class DegradationAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.params = cfg.get('PARAMS', {
+ 'gaussian_noise': {},
+ 'resize': {},
+ 'jpeg': {},
+ 'gaussian_blur': {},
+ })
+ if not isinstance(self.params, dict):
+ self.params = Config.get_dict(self.params)
+ self.random_degradation = cfg.get('RANDOM_DEGRADATION', False)
+
+ def forward(self, image):
+ if isinstance(image, Image.Image):
+ image = np.array(image)
+ elif isinstance(image, torch.Tensor):
+ image = image.detach().cpu().numpy()
+ elif isinstance(image, np.ndarray):
+ image = image.copy()
+ else:
+ raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+ if np.max(image) > 1.0:
+ image = (image / 255.).astype(np.float32)
+
+ degradation_list = list(self.params.keys())
+ if self.random_degradation:
+ random.shuffle(degradation_list)
+
+ for degradation_type in degradation_list:
+ if degradation_type == 'gaussian_noise':
+ image = gaussian_noise_op(image, self.params[degradation_type])
+ elif degradation_type == 'resize':
+ image = resize_op(image, self.params[degradation_type])
+ elif degradation_type == 'jpeg':
+ image = jpeg_op(image, self.params[degradation_type])
+ elif degradation_type == 'gaussian_blur':
+ image = gaussian_blur_op(image, self.params[degradation_type])
+ else:
+ raise NotImplementedError(
+ f'ERROR: degradation_type: {degradation_type} is invalid.')
+ image = (image * 255.0).astype(np.uint8)
+
+ assert len(image.shape) < 4
+ return image
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ DegradationAnnotator.para_dict,
+ set_name=True)
diff --git a/scepter/modules/annotator/doodle.py b/scepter/modules/annotator/doodle.py
new file mode 100644
index 0000000..89ec0a8
--- /dev/null
+++ b/scepter/modules/annotator/doodle.py
@@ -0,0 +1,51 @@
+# -*- coding: utf-8 -*-
+
+import math
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torchvision.transforms as TT
+from einops import rearrange
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+
+
+@ANNOTATORS.register_class()
+class DoodleAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.processor_type = cfg.get('PROCESSOR_TYPE', 'pidinet_sketch')
+ processor_cfg = cfg.get('PROCESSOR_CFG', None)
+ if self.processor_type == 'pidinet_sketch':
+ self.pidinet_ins = ANNOTATORS.build(processor_cfg[0])
+ self.sketch_ins = ANNOTATORS.build(processor_cfg[1])
+ else:
+ raise 'Unsurpport PROCESSOR for DoodleAnnotator'
+
+ @torch.no_grad()
+ @torch.inference_mode()
+ @torch.autocast('cuda', enabled=False)
+ def forward(self, image):
+ if self.processor_type == 'pidinet_sketch':
+ pidinet_res = self.pidinet_ins(image)
+ sketch_res = self.sketch_ins(pidinet_res)
+ doodle_res = sketch_res
+ else:
+ raise 'Unsurpport PROCESSOR for DoodleAnnotator'
+ return doodle_res
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ DoodleAnnotator.para_dict,
+ set_name=True)
diff --git a/scepter/modules/annotator/gray.py b/scepter/modules/annotator/gray.py
new file mode 100644
index 0000000..b718bdb
--- /dev/null
+++ b/scepter/modules/annotator/gray.py
@@ -0,0 +1,38 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+from PIL import Image
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import dict_to_yaml
+
+
+@ANNOTATORS.register_class()
+class GrayAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+
+ def forward(self, image):
+ if isinstance(image, Image.Image):
+ image = np.array(image)
+ elif isinstance(image, torch.Tensor):
+ image = image.detach().cpu().numpy()
+ elif isinstance(image, np.ndarray):
+ image = image.copy()
+ else:
+ raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+ gray_map = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
+ return gray_map[..., None].repeat(3, axis=2)
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ GrayAnnotator.para_dict,
+ set_name=True)
diff --git a/scepter/modules/annotator/informative_drawing.py b/scepter/modules/annotator/informative_drawing.py
new file mode 100644
index 0000000..0df8526
--- /dev/null
+++ b/scepter/modules/annotator/informative_drawing.py
@@ -0,0 +1,178 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+import torch.nn as nn
+import torchvision
+import torchvision.transforms as transforms
+from einops import rearrange
+from PIL import Image
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+from torchvision.transforms import InterpolationMode
+
+norm_layer = nn.InstanceNorm2d
+
+
+class ResidualBlock(nn.Module):
+ def __init__(self, in_features):
+ super(ResidualBlock, self).__init__()
+
+ conv_block = [
+ nn.ReflectionPad2d(1),
+ nn.Conv2d(in_features, in_features, 3),
+ norm_layer(in_features),
+ nn.ReLU(inplace=True),
+ nn.ReflectionPad2d(1),
+ nn.Conv2d(in_features, in_features, 3),
+ norm_layer(in_features)
+ ]
+
+ self.conv_block = nn.Sequential(*conv_block)
+
+ def forward(self, x):
+ return x + self.conv_block(x)
+
+
+class ContourInference(nn.Module):
+ def __init__(self, input_nc, output_nc, n_residual_blocks=9, sigmoid=True):
+ super(ContourInference, self).__init__()
+
+ # Initial convolution block
+ model0 = [
+ nn.ReflectionPad2d(3),
+ nn.Conv2d(input_nc, 64, 7),
+ norm_layer(64),
+ nn.ReLU(inplace=True)
+ ]
+ self.model0 = nn.Sequential(*model0)
+
+ # Downsampling
+ model1 = []
+ in_features = 64
+ out_features = in_features * 2
+ for _ in range(2):
+ model1 += [
+ nn.Conv2d(in_features, out_features, 3, stride=2, padding=1),
+ norm_layer(out_features),
+ nn.ReLU(inplace=True)
+ ]
+ in_features = out_features
+ out_features = in_features * 2
+ self.model1 = nn.Sequential(*model1)
+
+ model2 = []
+ # Residual blocks
+ for _ in range(n_residual_blocks):
+ model2 += [ResidualBlock(in_features)]
+ self.model2 = nn.Sequential(*model2)
+
+ # Upsampling
+ model3 = []
+ out_features = in_features // 2
+ for _ in range(2):
+ model3 += [
+ nn.ConvTranspose2d(in_features,
+ out_features,
+ 3,
+ stride=2,
+ padding=1,
+ output_padding=1),
+ norm_layer(out_features),
+ nn.ReLU(inplace=True)
+ ]
+ in_features = out_features
+ out_features = in_features // 2
+ self.model3 = nn.Sequential(*model3)
+
+ # Output layer
+ model4 = [nn.ReflectionPad2d(3), nn.Conv2d(64, output_nc, 7)]
+ if sigmoid:
+ model4 += [nn.Sigmoid()]
+
+ self.model4 = nn.Sequential(*model4)
+
+ def forward(self, x, cond=None):
+ out = self.model0(x)
+ out = self.model1(out)
+ out = self.model2(out)
+ out = self.model3(out)
+ out = self.model4(out)
+
+ return out
+
+
+@ANNOTATORS.register_class()
+class InfoDrawContourAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ input_nc = cfg.get('INPUT_NC', 3)
+ output_nc = cfg.get('OUTPUT_NC', 1)
+ n_residual_blocks = cfg.get('N_RESIDUAL_BLOCKS', 3)
+ sigmoid = cfg.get('SIGMOID', True)
+ pretrained_model = cfg.get('PRETRAINED_MODEL', None)
+
+ self.model = ContourInference(input_nc, output_nc, n_residual_blocks,
+ sigmoid)
+ with FS.get_from(pretrained_model, wait_finish=True) as local_path:
+ self.model.load_state_dict(torch.load(local_path))
+ self.model = self.model.eval().requires_grad_(False).to(we.device_id)
+
+ @torch.no_grad()
+ @torch.inference_mode()
+ @torch.autocast('cuda', enabled=False)
+ def forward(self, image):
+ is_batch = False if len(image.shape) == 3 else True
+ if isinstance(image, torch.Tensor):
+ if len(image.shape) == 3:
+ image = rearrange(image, 'h w c -> 1 c h w')
+ B, C, H, W = image.shape
+ elif len(image.shape) == 4:
+ B, C, H, W = image.shape
+ else:
+ raise "Unsurpport input image's shape"
+ elif isinstance(image, np.ndarray):
+ image = torch.from_numpy(image.copy()).float()
+ if len(image.shape) == 3:
+ image = rearrange(image, 'h w c -> 1 c h w')
+ B, C, H, W = image.shape
+ elif len(image.shape) == 4:
+ B, C, H, W = image.shape
+ else:
+ raise "Unsurpport input image's shape"
+ else:
+ raise "Unsurpport input image's type"
+
+ image = image.float().div(255).to(we.device_id)
+ contour_map = self.model(image)
+ contour_map = (contour_map.squeeze(dim=1) * 255.0).clip(
+ 0, 255).cpu().numpy().astype(np.uint8)
+ contour_map = contour_map[..., None].repeat(3, -1)
+ if not is_batch:
+ contour_map = contour_map.squeeze()
+ return contour_map
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ InfoDrawContourAnnotator.para_dict,
+ set_name=True)
+
+
+@ANNOTATORS.register_class()
+class InfoDrawAnimeAnnotator(InfoDrawContourAnnotator):
+ pass
+
+
+@ANNOTATORS.register_class()
+class InfoDrawOpenSketchAnnotator(InfoDrawContourAnnotator):
+ pass
diff --git a/scepter/modules/annotator/inpainting.py b/scepter/modules/annotator/inpainting.py
new file mode 100644
index 0000000..2f422f1
--- /dev/null
+++ b/scepter/modules/annotator/inpainting.py
@@ -0,0 +1,271 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import math
+import random
+from abc import ABCMeta
+from enum import Enum
+
+import cv2
+import numpy as np
+import torch
+from PIL import Image, ImageDraw
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import Config, dict_to_yaml
+
+
+def invert_image(im):
+ im_arr = np.array(im)
+ mask_1 = im_arr == 0
+ mask_2 = im_arr == 255
+ im_arr[mask_1] = 255
+ im_arr[mask_2] = 0
+ new_im = im_arr
+ return new_im
+
+
+class DrawMethod(Enum):
+ LINE = 'line'
+ CIRCLE = 'circle'
+ SQUARE = 'square'
+
+
+def make_random_irregular_mask(shape,
+ max_angle=4,
+ max_len=60,
+ max_width=20,
+ min_times=0,
+ max_times=10,
+ draw_method=DrawMethod.LINE):
+ draw_method = DrawMethod(draw_method)
+
+ height, width = shape
+ mask = np.zeros((height, width), np.float32)
+ times = np.random.randint(min_times, max_times + 1)
+ for i in range(times):
+ start_x = np.random.randint(width)
+ start_y = np.random.randint(height)
+ for j in range(1 + np.random.randint(5)):
+ angle = 0.01 + np.random.randint(max_angle)
+ if i % 2 == 0:
+ angle = 2 * 3.1415926 - angle
+ length = 10 + np.random.randint(max_len)
+ brush_w = 5 + np.random.randint(max_width)
+ end_x = np.clip(
+ (start_x + length * np.sin(angle)).astype(np.int32), 0, width)
+ end_y = np.clip(
+ (start_y + length * np.cos(angle)).astype(np.int32), 0, height)
+ if draw_method == DrawMethod.LINE:
+ cv2.line(mask, (start_x, start_y), (end_x, end_y), 1.0,
+ brush_w)
+ elif draw_method == DrawMethod.CIRCLE:
+ cv2.circle(mask, (start_x, start_y),
+ radius=brush_w,
+ color=1.,
+ thickness=-1)
+ elif draw_method == DrawMethod.SQUARE:
+ radius = brush_w // 2
+ mask[start_y - radius:start_y + radius,
+ start_x - radius:start_x + radius] = 1
+ start_x, start_y = end_x, end_y
+ return mask[None, ...]
+
+
+class RandomIrregularMaskGenerator:
+ def __init__(self,
+ max_angle=4,
+ max_len=60,
+ max_width=20,
+ min_times=0,
+ max_times=10,
+ ramp_kwargs=None,
+ draw_method=DrawMethod.LINE):
+ self.max_angle = max_angle
+ self.max_len = max_len
+ self.max_width = max_width
+ self.min_times = min_times
+ self.max_times = max_times
+ self.draw_method = draw_method
+ self.ramp = None
+ # self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
+
+ def __call__(self, img, iter_i=None, raw_image=None):
+ coef = self.ramp(iter_i) if (self.ramp is not None) and (
+ iter_i is not None) else 1
+ cur_max_len = int(max(1, self.max_len * coef))
+ cur_max_width = int(max(1, self.max_width * coef))
+ cur_max_times = int(self.min_times + 1 +
+ (self.max_times - self.min_times) * coef)
+ return make_random_irregular_mask(img.shape[1:],
+ max_angle=self.max_angle,
+ max_len=cur_max_len,
+ max_width=cur_max_width,
+ min_times=self.min_times,
+ max_times=cur_max_times,
+ draw_method=self.draw_method)
+
+
+def make_random_rectangle_mask(shape,
+ margin=10,
+ bbox_min_size=30,
+ bbox_max_size=100,
+ min_times=0,
+ max_times=3):
+ height, width = shape
+ mask = np.zeros((height, width), np.float32)
+ bbox_max_size = min(bbox_max_size, height - margin * 2, width - margin * 2)
+ times = np.random.randint(min_times, max_times + 1)
+ for i in range(times):
+ box_width = np.random.randint(bbox_min_size, bbox_max_size)
+ box_height = np.random.randint(bbox_min_size, bbox_max_size)
+ start_x = np.random.randint(margin, width - margin - box_width + 1)
+ start_y = np.random.randint(margin, height - margin - box_height + 1)
+ mask[start_y:start_y + box_height, start_x:start_x + box_width] = 1
+ return mask[None, ...]
+
+
+class RandomRectangleMaskGenerator:
+ def __init__(self,
+ margin=10,
+ bbox_min_size=30,
+ bbox_max_size=100,
+ min_times=0,
+ max_times=3,
+ ramp_kwargs=None):
+ self.margin = margin
+ self.bbox_min_size = bbox_min_size
+ self.bbox_max_size = bbox_max_size
+ self.min_times = min_times
+ self.max_times = max_times
+ self.ramp = None
+ # self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
+
+ def __call__(self, img, iter_i=None, raw_image=None):
+ coef = self.ramp(iter_i) if (self.ramp is not None) and (
+ iter_i is not None) else 1
+ cur_bbox_max_size = int(self.bbox_min_size + 1 +
+ (self.bbox_max_size - self.bbox_min_size) *
+ coef)
+ cur_max_times = int(self.min_times +
+ (self.max_times - self.min_times) * coef)
+ return make_random_rectangle_mask(img.shape[1:],
+ margin=self.margin,
+ bbox_min_size=self.bbox_min_size,
+ bbox_max_size=cur_bbox_max_size,
+ min_times=self.min_times,
+ max_times=cur_max_times)
+
+
+class MixedMaskGenerator:
+ def __init__(self,
+ irregular_proba=1 / 3,
+ irregular_kwargs=None,
+ box_proba=1 / 3,
+ box_kwargs=None,
+ invert_proba=0):
+ self.probas = []
+ self.gens = []
+
+ if irregular_proba > 0:
+ self.probas.append(irregular_proba)
+ if irregular_kwargs is None:
+ irregular_kwargs = {}
+ else:
+ irregular_kwargs = dict(irregular_kwargs)
+ irregular_kwargs['draw_method'] = DrawMethod.LINE
+ self.gens.append(RandomIrregularMaskGenerator(**irregular_kwargs))
+
+ if box_proba > 0:
+ self.probas.append(box_proba)
+ if box_kwargs is None:
+ box_kwargs = {}
+ self.gens.append(RandomRectangleMaskGenerator(**box_kwargs))
+
+ self.probas = np.array(self.probas, dtype='float32')
+ self.probas /= self.probas.sum()
+ self.invert_proba = invert_proba
+
+ def __call__(self, img, iter_i=None, raw_image=None):
+ kind = np.random.choice(len(self.probas), p=self.probas)
+ gen = self.gens[kind]
+ result = gen(img, iter_i=iter_i, raw_image=raw_image)
+ if self.invert_proba > 0 and random.random() < self.invert_proba:
+ result = 1 - result
+ return result
+
+
+@ANNOTATORS.register_class()
+class InpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.mask_cfg = cfg.get(
+ 'MASK_CFG', {
+ 'irregular_proba': 0.5,
+ 'irregular_kwargs': {
+ 'min_times': 4,
+ 'max_times': 10,
+ 'max_width': 100,
+ 'max_angle': 4,
+ 'max_len': 200
+ },
+ 'box_proba': 0.5,
+ 'box_kwargs': {
+ 'margin': 0,
+ 'bbox_min_size': 30,
+ 'bbox_max_size': 150,
+ 'max_times': 5,
+ 'min_times': 1
+ }
+ })
+ self.mask_cfg = Config.get_dict(self.mask_cfg) if not isinstance(
+ self.mask_cfg, dict) else self.mask_cfg
+ self.mask_generator = MixedMaskGenerator(**self.mask_cfg)
+ self.return_mask = cfg.get('RETURN_MASK', False)
+ self.return_invert = cfg.get('RETURN_INVERT', True)
+ self.mask_color = cfg.get('MASK_COLOR', 0)
+
+ def forward(self, image, mask=None, return_mask=None, return_invert=None, mask_color=None):
+ return_mask = return_mask if return_mask is not None else self.return_mask
+ return_invert = return_invert if return_invert is not None else self.return_invert
+ mask_color = mask_color if mask_color is not None else self.mask_color
+ if isinstance(image, Image.Image):
+ image = np.array(image)
+ elif isinstance(image, torch.Tensor):
+ image = image.detach().cpu().numpy()
+ elif isinstance(image, np.ndarray):
+ image = image.copy()
+ else:
+ raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+
+ if mask is not None:
+ if mask_color:
+ image[np.array(mask) == 255] = mask_color
+ else:
+ image[np.array(mask) == 255] = 0
+ else:
+ img = np.transpose(image, (2, 0, 1))
+ mask = self.mask_generator(img)
+ mask = (np.transpose(mask, (1, 2, 0)).squeeze(-1) * 255).astype(np.uint8)
+ if return_invert:
+ mask = invert_image(mask)
+ colored_mask = np.zeros_like(image)
+ if mask_color: colored_mask[:] = mask_color
+ image = np.where(mask[:, :, np.newaxis] == 255, colored_mask, image)
+
+ if return_mask:
+ ret_data = {
+ 'image': np.array(image),
+ 'mask': np.array(mask)
+ }
+ else:
+ ret_data = np.array(image)
+ return ret_data
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ InpaintingAnnotator.para_dict,
+ set_name=True)
diff --git a/scepter/modules/annotator/lama.py b/scepter/modules/annotator/lama.py
new file mode 100644
index 0000000..c0a6626
--- /dev/null
+++ b/scepter/modules/annotator/lama.py
@@ -0,0 +1,98 @@
+from abc import ABCMeta
+
+import torch
+import cv2
+import numpy as np
+from PIL import Image
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+
+def dilate_mask(mask, dilate_factor=15):
+ mask = mask.astype(np.uint8)
+ mask = cv2.dilate(
+ mask,
+ np.ones((dilate_factor, dilate_factor), np.uint8),
+ iterations=1
+ )
+ return mask
+
+@ANNOTATORS.register_class()
+class LamaAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ from modelscope.pipelines.builder import PIPELINES
+ from modelscope.pipelines.cv import ImageInpaintingPipeline
+ from modelscope.pipelines import pipeline
+ from modelscope.utils.constant import Tasks
+ from modelscope.metainfo import Pipelines
+ from modelscope.models.cv.image_inpainting.refinement import refine_predict
+ from torch.utils.data._utils.collate import default_collate
+
+ @PIPELINES.register_module(Tasks.image_inpainting, module_name=Pipelines.image_inpainting + "-v2")
+ class ImageInpaintingPipelineV2(ImageInpaintingPipeline):
+ def perform_inference(self, data):
+ px_budget = 9000000
+ batch = default_collate([data])
+ if self.refine:
+ assert 'unpad_to_size' in batch, 'Unpadded size is required for the refinement'
+ assert 'cuda' in str(self.device), 'GPU is required for refinement'
+ gpu_ids = str(self.device).split(':')[-1]
+ cur_res = refine_predict(
+ batch,
+ self.infer_model,
+ gpu_ids=gpu_ids,
+ modulo=self.pad_out_to_modulo,
+ n_iters=15,
+ lr=0.002,
+ min_side=512,
+ max_scales=3,
+ px_budget=px_budget)
+ cur_res = cur_res[0].permute(1, 2, 0).detach().cpu().numpy()
+ else:
+ with torch.no_grad():
+ batch = self.move_to_device(batch, self.device)
+ batch['mask'] = (batch['mask'] > 0) * 1
+ batch = self.infer_model(batch)
+ cur_res = batch['inpainted'][0].permute(
+ 1, 2, 0).detach().cpu().numpy()
+ unpad_to_size = batch.get('unpad_to_size', None)
+ if unpad_to_size is not None:
+ orig_height, orig_width = unpad_to_size
+ cur_res = cur_res[:orig_height, :orig_width]
+
+ cur_res = np.clip(cur_res * 255, 0, 255).astype('uint8')
+ cur_res = cv2.cvtColor(cur_res, cv2.COLOR_RGB2BGR)
+ return cur_res
+
+ lama_model_dir = FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL)
+ self.lama_model = pipeline(Tasks.image_inpainting, model=lama_model_dir,
+ pipeline_name=Pipelines.image_inpainting + "-v2", refine=True,
+ device="cuda:{}".format(we.device_id))
+ def forward(self, image, mask):
+ mask = dilate_mask(mask, dilate_factor=19)
+ input_mask = Image.fromarray(mask)
+ mask_expanded = np.tile(np.expand_dims(mask, axis=-1), (1, 1, 3))
+ input_image_np = np.array(image)
+ input_image_np[mask_expanded == 255] = 0
+ input_image = Image.fromarray(input_image_np)
+ input = {
+ 'img': input_image,
+ 'mask': input_mask,
+ }
+ result = self.lama_model(input)
+ output_img = result['output_img']
+ return output_img[..., ::-1]
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ LamaAnnotator.para_dict,
+ set_name=True)
+
+
+
diff --git a/scepter/modules/annotator/openpose.py b/scepter/modules/annotator/openpose.py
index b95fd9e..db4a76f 100644
--- a/scepter/modules/annotator/openpose.py
+++ b/scepter/modules/annotator/openpose.py
@@ -9,9 +9,10 @@ import os
from abc import ABCMeta
from collections import OrderedDict
+import numpy as np
+
import cv2
import matplotlib
-import numpy as np
import torch
import torch.nn as nn
from PIL import Image
@@ -478,7 +479,7 @@ class Hand(object):
map_ori = heatmap_avg[:, :, part]
one_heatmap = gaussian_filter(map_ori, sigma=3)
binary = np.ascontiguousarray(one_heatmap > thre, dtype=np.uint8)
- # 全部小于阈值
+
if np.sum(binary) == 0:
all_peaks.append([0, 0])
continue
diff --git a/scepter/modules/annotator/outpainting.py b/scepter/modules/annotator/outpainting.py
new file mode 100644
index 0000000..f1686a5
--- /dev/null
+++ b/scepter/modules/annotator/outpainting.py
@@ -0,0 +1,189 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import math
+import random
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+from PIL import Image, ImageDraw
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import dict_to_yaml
+
+
+@ANNOTATORS.register_class()
+class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.mask_blur = cfg.get('MASK_BLUR', 0)
+ self.random_cfg = cfg.get('RANDOM_CFG', None)
+ self.return_mask = cfg.get('RETURN_MASK', False)
+ self.return_source = cfg.get('RETURN_SOURCE', True)
+ self.keep_padding_ratio = cfg.get('KEEP_PADDING_RATIO', 64)
+ self.mask_color = cfg.get('MASK_COLOR', 0)
+
+ def get_box(self, mask):
+ locs = np.where(mask == 255)
+ if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
+ return None
+ left, right = np.min(locs[1]), np.max(locs[1])
+ top, bottom = np.min(locs[0]), np.max(locs[0])
+ return [left, top, right, bottom]
+
+ def forward(self,
+ image,
+ ratio=0.3,
+ mask=None,
+ direction=['left', 'right', 'up', 'down'],
+ return_mask=None,
+ return_source=None,
+ mask_color=None):
+ return_mask = return_mask if return_mask is not None else self.return_mask
+ return_source = return_source if return_source is not None else self.return_source
+ mask_color = mask_color if mask_color is not None else self.mask_color
+ if isinstance(image, Image.Image):
+ image = image
+ elif isinstance(image, torch.Tensor):
+ image = Image.fromarray(image.detach().cpu().numpy())
+ elif isinstance(image, np.ndarray):
+ image = Image.fromarray(image.copy())
+ else:
+ raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+
+ if self.random_cfg:
+ direction_range = self.random_cfg.get(
+ 'DIRECTION_RANGE', ['left', 'right', 'up', 'down'])
+ ratio_range = self.random_cfg.get('RATIO_RANGE', [0.0, 1.0])
+ direction = random.sample(
+ direction_range,
+ random.choice(list(range(1,
+ len(direction_range) + 1))))
+ ratio = random.uniform(ratio_range[0], ratio_range[1])
+
+ if mask is None:
+ init_image = image
+ src_width, src_height = init_image.width, init_image.height
+ left = int(ratio * src_width) if 'left' in direction else 0
+ right = int(ratio * src_width) if 'right' in direction else 0
+ up = int(ratio * src_height) if 'up' in direction else 0
+ down = int(ratio * src_height) if 'down' in direction else 0
+ # print(direction, ratio, left, right, up, down)
+ tar_width = math.ceil(
+ (src_width + left + right) /
+ self.keep_padding_ratio) * self.keep_padding_ratio
+ tar_height = math.ceil(
+ (src_height + up + down) /
+ self.keep_padding_ratio) * self.keep_padding_ratio
+ if left > 0:
+ left = left * (tar_width - src_width) // (left + right)
+ if right > 0:
+ right = tar_width - src_width - left
+ if up > 0:
+ up = up * (tar_height - src_height) // (up + down)
+ if down > 0:
+ down = tar_height - src_height - up
+ if mask_color is not None:
+ img = Image.new('RGB', (tar_width, tar_height), color=mask_color)
+ else:
+ img = Image.new('RGB', (tar_width, tar_height))
+ img.paste(init_image, (left, up))
+ mask = Image.new('L', (img.width, img.height), 'white')
+ draw = ImageDraw.Draw(mask)
+
+ draw.rectangle(
+ (left + (self.mask_blur * 2 if left > 0 else 0), up +
+ (self.mask_blur * 2 if up > 0 else 0), mask.width - right -
+ (self.mask_blur * 2 if right > 0 else 0), mask.height - down -
+ (self.mask_blur * 2 if down > 0 else 0)),
+ fill='black')
+ else:
+ bbox = self.get_box(np.array(mask))
+ if bbox is None:
+ img = image
+ mask = mask
+ init_image = image
+ else:
+ mask = Image.new('L', (image.width, image.height), 'white')
+ mask_zero = Image.new('L', (bbox[2]-bbox[0], bbox[3]-bbox[1]), 'black')
+ mask.paste(mask_zero, (bbox[0], bbox[1]))
+ crop_image = image.crop(bbox)
+ init_image = Image.new('RGB', (image.width, image.height), 'black')
+ init_image.paste(crop_image, (bbox[0], bbox[1]))
+ img = image
+ if return_mask:
+ if return_source:
+ ret_data = {'src_image': np.array(init_image), 'image': np.array(img), 'mask': np.array(mask)}
+ else:
+ ret_data = {'image': np.array(img), 'mask': np.array(mask)}
+ else:
+ if return_source:
+ ret_data = {'src_image': np.array(init_image), 'image': np.array(img)}
+ else:
+ ret_data = np.array(img)
+ return ret_data
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ OutpaintingAnnotator.para_dict,
+ set_name=True)
+
+@ANNOTATORS.register_class()
+class OutpaintingResize(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+
+ def get_box(self, mask):
+ locs = np.where(mask == 0)
+ if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
+ return None
+ left, right = np.min(locs[1]), np.max(locs[1])
+ top, bottom = np.min(locs[0]), np.max(locs[0])
+ return [left, top, right, bottom]
+
+ def forward(self,
+ image,
+ target_image,
+ mask=None
+ ):
+ if isinstance(image, Image.Image):
+ image = image
+ elif isinstance(image, torch.Tensor):
+ image = Image.fromarray(image.detach().cpu().numpy())
+ elif isinstance(image, np.ndarray):
+ image = Image.fromarray(image.copy())
+ else:
+ raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+
+ if isinstance(target_image, Image.Image):
+ target_image = target_image
+ elif isinstance(target_image, torch.Tensor):
+ target_image = Image.fromarray(target_image.detach().cpu().numpy())
+ elif isinstance(target_image, np.ndarray):
+ target_image = Image.fromarray(target_image.copy())
+ else:
+ raise f'Unsurpport datatype{type(target_image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+
+ bbox = self.get_box(np.array(mask))
+ if bbox is None:
+ init_image = image
+ else:
+ paste_img = image.resize((bbox[2]-bbox[0], bbox[3]-bbox[1]))
+ init_image = Image.new('RGB', (target_image.width, target_image.height), 'black')
+ init_image.paste(paste_img, (bbox[0], bbox[1]))
+ ret_data = {'src_image': np.array(init_image)}
+ return ret_data
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ OutpaintingResize.para_dict,
+ set_name=True)
diff --git a/scepter/modules/annotator/pidinet.py b/scepter/modules/annotator/pidinet.py
new file mode 100644
index 0000000..06a4bc8
--- /dev/null
+++ b/scepter/modules/annotator/pidinet.py
@@ -0,0 +1,935 @@
+# -*- coding: utf-8 -*-
+
+import math
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torchvision.transforms as TT
+from einops import rearrange
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+
+CONFIGS = {
+ 'baseline': {
+ 'layer0': 'cv',
+ 'layer1': 'cv',
+ 'layer2': 'cv',
+ 'layer3': 'cv',
+ 'layer4': 'cv',
+ 'layer5': 'cv',
+ 'layer6': 'cv',
+ 'layer7': 'cv',
+ 'layer8': 'cv',
+ 'layer9': 'cv',
+ 'layer10': 'cv',
+ 'layer11': 'cv',
+ 'layer12': 'cv',
+ 'layer13': 'cv',
+ 'layer14': 'cv',
+ 'layer15': 'cv',
+ },
+ 'c-v15': {
+ 'layer0': 'cd',
+ 'layer1': 'cv',
+ 'layer2': 'cv',
+ 'layer3': 'cv',
+ 'layer4': 'cv',
+ 'layer5': 'cv',
+ 'layer6': 'cv',
+ 'layer7': 'cv',
+ 'layer8': 'cv',
+ 'layer9': 'cv',
+ 'layer10': 'cv',
+ 'layer11': 'cv',
+ 'layer12': 'cv',
+ 'layer13': 'cv',
+ 'layer14': 'cv',
+ 'layer15': 'cv',
+ },
+ 'a-v15': {
+ 'layer0': 'ad',
+ 'layer1': 'cv',
+ 'layer2': 'cv',
+ 'layer3': 'cv',
+ 'layer4': 'cv',
+ 'layer5': 'cv',
+ 'layer6': 'cv',
+ 'layer7': 'cv',
+ 'layer8': 'cv',
+ 'layer9': 'cv',
+ 'layer10': 'cv',
+ 'layer11': 'cv',
+ 'layer12': 'cv',
+ 'layer13': 'cv',
+ 'layer14': 'cv',
+ 'layer15': 'cv',
+ },
+ 'r-v15': {
+ 'layer0': 'rd',
+ 'layer1': 'cv',
+ 'layer2': 'cv',
+ 'layer3': 'cv',
+ 'layer4': 'cv',
+ 'layer5': 'cv',
+ 'layer6': 'cv',
+ 'layer7': 'cv',
+ 'layer8': 'cv',
+ 'layer9': 'cv',
+ 'layer10': 'cv',
+ 'layer11': 'cv',
+ 'layer12': 'cv',
+ 'layer13': 'cv',
+ 'layer14': 'cv',
+ 'layer15': 'cv',
+ },
+ 'cvvv4': {
+ 'layer0': 'cd',
+ 'layer1': 'cv',
+ 'layer2': 'cv',
+ 'layer3': 'cv',
+ 'layer4': 'cd',
+ 'layer5': 'cv',
+ 'layer6': 'cv',
+ 'layer7': 'cv',
+ 'layer8': 'cd',
+ 'layer9': 'cv',
+ 'layer10': 'cv',
+ 'layer11': 'cv',
+ 'layer12': 'cd',
+ 'layer13': 'cv',
+ 'layer14': 'cv',
+ 'layer15': 'cv',
+ },
+ 'avvv4': {
+ 'layer0': 'ad',
+ 'layer1': 'cv',
+ 'layer2': 'cv',
+ 'layer3': 'cv',
+ 'layer4': 'ad',
+ 'layer5': 'cv',
+ 'layer6': 'cv',
+ 'layer7': 'cv',
+ 'layer8': 'ad',
+ 'layer9': 'cv',
+ 'layer10': 'cv',
+ 'layer11': 'cv',
+ 'layer12': 'ad',
+ 'layer13': 'cv',
+ 'layer14': 'cv',
+ 'layer15': 'cv',
+ },
+ 'rvvv4': {
+ 'layer0': 'rd',
+ 'layer1': 'cv',
+ 'layer2': 'cv',
+ 'layer3': 'cv',
+ 'layer4': 'rd',
+ 'layer5': 'cv',
+ 'layer6': 'cv',
+ 'layer7': 'cv',
+ 'layer8': 'rd',
+ 'layer9': 'cv',
+ 'layer10': 'cv',
+ 'layer11': 'cv',
+ 'layer12': 'rd',
+ 'layer13': 'cv',
+ 'layer14': 'cv',
+ 'layer15': 'cv',
+ },
+ 'cccv4': {
+ 'layer0': 'cd',
+ 'layer1': 'cd',
+ 'layer2': 'cd',
+ 'layer3': 'cv',
+ 'layer4': 'cd',
+ 'layer5': 'cd',
+ 'layer6': 'cd',
+ 'layer7': 'cv',
+ 'layer8': 'cd',
+ 'layer9': 'cd',
+ 'layer10': 'cd',
+ 'layer11': 'cv',
+ 'layer12': 'cd',
+ 'layer13': 'cd',
+ 'layer14': 'cd',
+ 'layer15': 'cv',
+ },
+ 'aaav4': {
+ 'layer0': 'ad',
+ 'layer1': 'ad',
+ 'layer2': 'ad',
+ 'layer3': 'cv',
+ 'layer4': 'ad',
+ 'layer5': 'ad',
+ 'layer6': 'ad',
+ 'layer7': 'cv',
+ 'layer8': 'ad',
+ 'layer9': 'ad',
+ 'layer10': 'ad',
+ 'layer11': 'cv',
+ 'layer12': 'ad',
+ 'layer13': 'ad',
+ 'layer14': 'ad',
+ 'layer15': 'cv',
+ },
+ 'rrrv4': {
+ 'layer0': 'rd',
+ 'layer1': 'rd',
+ 'layer2': 'rd',
+ 'layer3': 'cv',
+ 'layer4': 'rd',
+ 'layer5': 'rd',
+ 'layer6': 'rd',
+ 'layer7': 'cv',
+ 'layer8': 'rd',
+ 'layer9': 'rd',
+ 'layer10': 'rd',
+ 'layer11': 'cv',
+ 'layer12': 'rd',
+ 'layer13': 'rd',
+ 'layer14': 'rd',
+ 'layer15': 'cv',
+ },
+ 'c16': {
+ 'layer0': 'cd',
+ 'layer1': 'cd',
+ 'layer2': 'cd',
+ 'layer3': 'cd',
+ 'layer4': 'cd',
+ 'layer5': 'cd',
+ 'layer6': 'cd',
+ 'layer7': 'cd',
+ 'layer8': 'cd',
+ 'layer9': 'cd',
+ 'layer10': 'cd',
+ 'layer11': 'cd',
+ 'layer12': 'cd',
+ 'layer13': 'cd',
+ 'layer14': 'cd',
+ 'layer15': 'cd',
+ },
+ 'a16': {
+ 'layer0': 'ad',
+ 'layer1': 'ad',
+ 'layer2': 'ad',
+ 'layer3': 'ad',
+ 'layer4': 'ad',
+ 'layer5': 'ad',
+ 'layer6': 'ad',
+ 'layer7': 'ad',
+ 'layer8': 'ad',
+ 'layer9': 'ad',
+ 'layer10': 'ad',
+ 'layer11': 'ad',
+ 'layer12': 'ad',
+ 'layer13': 'ad',
+ 'layer14': 'ad',
+ 'layer15': 'ad',
+ },
+ 'r16': {
+ 'layer0': 'rd',
+ 'layer1': 'rd',
+ 'layer2': 'rd',
+ 'layer3': 'rd',
+ 'layer4': 'rd',
+ 'layer5': 'rd',
+ 'layer6': 'rd',
+ 'layer7': 'rd',
+ 'layer8': 'rd',
+ 'layer9': 'rd',
+ 'layer10': 'rd',
+ 'layer11': 'rd',
+ 'layer12': 'rd',
+ 'layer13': 'rd',
+ 'layer14': 'rd',
+ 'layer15': 'rd',
+ },
+ 'carv4': {
+ 'layer0': 'cd',
+ 'layer1': 'ad',
+ 'layer2': 'rd',
+ 'layer3': 'cv',
+ 'layer4': 'cd',
+ 'layer5': 'ad',
+ 'layer6': 'rd',
+ 'layer7': 'cv',
+ 'layer8': 'cd',
+ 'layer9': 'ad',
+ 'layer10': 'rd',
+ 'layer11': 'cv',
+ 'layer12': 'cd',
+ 'layer13': 'ad',
+ 'layer14': 'rd',
+ 'layer15': 'cv'
+ }
+}
+
+
+def create_conv_func(op_type):
+ assert op_type in ['cv', 'cd', 'ad',
+ 'rd'], 'unknown op type: %s' % str(op_type)
+ if op_type == 'cv':
+ return F.conv2d
+ if op_type == 'cd':
+
+ def func(x,
+ weights,
+ bias=None,
+ stride=1,
+ padding=0,
+ dilation=1,
+ groups=1):
+ assert dilation in [1,
+ 2], 'dilation for cd_conv should be in 1 or 2'
+ assert weights.size(2) == 3 and weights.size(3) == 3, \
+ 'kernel size for cd_conv should be 3x3'
+ assert padding == dilation, 'padding for cd_conv set wrong'
+
+ weights_c = weights.sum(dim=[2, 3], keepdim=True)
+ yc = F.conv2d(x,
+ weights_c,
+ stride=stride,
+ padding=0,
+ groups=groups)
+ y = F.conv2d(x,
+ weights,
+ bias,
+ stride=stride,
+ padding=padding,
+ dilation=dilation,
+ groups=groups)
+ return y - yc
+
+ return func
+ elif op_type == 'ad':
+
+ def func(x,
+ weights,
+ bias=None,
+ stride=1,
+ padding=0,
+ dilation=1,
+ groups=1):
+ assert dilation in [1,
+ 2], 'dilation for ad_conv should be in 1 or 2'
+ assert weights.size(2) == 3 and weights.size(3) == 3, \
+ 'kernel size for ad_conv should be 3x3'
+ assert padding == dilation, 'padding for ad_conv set wrong'
+
+ shape = weights.shape
+ weights = weights.view(shape[0], shape[1], -1)
+ # clock-wise
+ weights_conv = (
+ weights -
+ weights[:, :, [3, 0, 1, 6, 4, 2, 7, 8, 5]]).view(shape)
+ y = F.conv2d(x,
+ weights_conv,
+ bias,
+ stride=stride,
+ padding=padding,
+ dilation=dilation,
+ groups=groups)
+ return y
+
+ return func
+ elif op_type == 'rd':
+
+ def func(x,
+ weights,
+ bias=None,
+ stride=1,
+ padding=0,
+ dilation=1,
+ groups=1):
+ assert dilation in [1,
+ 2], 'dilation for rd_conv should be in 1 or 2'
+ assert weights.size(2) == 3 and weights.size(3) == 3, \
+ 'kernel size for rd_conv should be 3x3'
+ padding = 2 * dilation
+
+ shape = weights.shape
+ if weights.is_cuda:
+ buffer = torch.cuda.FloatTensor(shape[0], shape[1],
+ 5 * 5).fill_(0)
+ else:
+ buffer = torch.zeros(shape[0], shape[1], 5 * 5)
+ weights = weights.view(shape[0], shape[1], -1)
+ buffer[:, :, [0, 2, 4, 10, 14, 20, 22, 24]] = weights[:, :, 1:]
+ buffer[:, :, [6, 7, 8, 11, 13, 16, 17, 18]] = -weights[:, :, 1:]
+ buffer[:, :, 12] = 0
+ buffer = buffer.view(shape[0], shape[1], 5, 5)
+ y = F.conv2d(x,
+ buffer,
+ bias,
+ stride=stride,
+ padding=padding,
+ dilation=dilation,
+ groups=groups)
+ return y
+
+ return func
+ else:
+ print('impossible to be here unless you force that', flush=True)
+ return None
+
+
+def config_model(model):
+ model_options = list(CONFIGS.keys())
+ assert model in model_options, \
+ 'unrecognized model, please choose from %s' % str(model_options)
+
+ pdcs = []
+ for i in range(16):
+ layer_name = 'layer%d' % i
+ op = CONFIGS[model][layer_name]
+ pdcs.append(create_conv_func(op))
+ return pdcs
+
+
+def config_model_converted(model):
+ model_options = list(CONFIGS.keys())
+ assert model in model_options, \
+ 'unrecognized model, please choose from %s' % str(model_options)
+
+ pdcs = []
+ for i in range(16):
+ layer_name = 'layer%d' % i
+ op = CONFIGS[model][layer_name]
+ pdcs.append(op)
+ return pdcs
+
+
+def convert_pdc(op, weight):
+ if op == 'cv':
+ return weight
+ elif op == 'cd':
+ shape = weight.shape
+ weight_c = weight.sum(dim=[2, 3])
+ weight = weight.view(shape[0], shape[1], -1)
+ weight[:, :, 4] = weight[:, :, 4] - weight_c
+ weight = weight.view(shape)
+ return weight
+ elif op == 'ad':
+ shape = weight.shape
+ weight = weight.view(shape[0], shape[1], -1)
+ weight_conv = (weight -
+ weight[:, :, [3, 0, 1, 6, 4, 2, 7, 8, 5]]).view(shape)
+ return weight_conv
+ elif op == 'rd':
+ shape = weight.shape
+ buffer = torch.zeros(shape[0], shape[1], 5 * 5, device=weight.device)
+ weight = weight.view(shape[0], shape[1], -1)
+ buffer[:, :, [0, 2, 4, 10, 14, 20, 22, 24]] = weight[:, :, 1:]
+ buffer[:, :, [6, 7, 8, 11, 13, 16, 17, 18]] = -weight[:, :, 1:]
+ buffer = buffer.view(shape[0], shape[1], 5, 5)
+ return buffer
+ raise ValueError('wrong op {}'.format(str(op)))
+
+
+def convert_pidinet(state_dict, config):
+ pdcs = config_model_converted(config)
+ new_dict = {}
+ for pname, p in state_dict.items():
+ if 'init_block.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[0], p)
+ elif 'block1_1.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[1], p)
+ elif 'block1_2.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[2], p)
+ elif 'block1_3.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[3], p)
+ elif 'block2_1.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[4], p)
+ elif 'block2_2.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[5], p)
+ elif 'block2_3.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[6], p)
+ elif 'block2_4.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[7], p)
+ elif 'block3_1.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[8], p)
+ elif 'block3_2.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[9], p)
+ elif 'block3_3.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[10], p)
+ elif 'block3_4.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[11], p)
+ elif 'block4_1.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[12], p)
+ elif 'block4_2.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[13], p)
+ elif 'block4_3.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[14], p)
+ elif 'block4_4.conv1.weight' in pname:
+ new_dict[pname] = convert_pdc(pdcs[15], p)
+ else:
+ new_dict[pname] = p
+ return new_dict
+
+
+class Conv2d(nn.Module):
+ def __init__(self,
+ pdc,
+ in_channels,
+ out_channels,
+ kernel_size,
+ stride=1,
+ padding=0,
+ dilation=1,
+ groups=1,
+ bias=False):
+ super().__init__()
+ if in_channels % groups != 0:
+ raise ValueError('in_channels must be divisible by groups')
+ if out_channels % groups != 0:
+ raise ValueError('out_channels must be divisible by groups')
+ self.in_channels = in_channels
+ self.out_channels = out_channels
+ self.kernel_size = kernel_size
+ self.stride = stride
+ self.padding = padding
+ self.dilation = dilation
+ self.groups = groups
+ self.weight = nn.Parameter(
+ torch.Tensor(out_channels, in_channels // groups, kernel_size,
+ kernel_size))
+ if bias:
+ self.bias = nn.Parameter(torch.Tensor(out_channels))
+ else:
+ self.register_parameter('bias', None)
+ self.reset_parameters()
+ self.pdc = pdc
+
+ def reset_parameters(self):
+ nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
+ if self.bias is not None:
+ fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
+ bound = 1 / math.sqrt(fan_in)
+ nn.init.uniform_(self.bias, -bound, bound)
+
+ def forward(self, input):
+ return self.pdc(input, self.weight, self.bias, self.stride,
+ self.padding, self.dilation, self.groups)
+
+
+class CSAM(nn.Module):
+ """
+ Compact Spatial Attention Module
+ """
+ def __init__(self, channels):
+ super().__init__()
+
+ mid_channels = 4
+ self.relu1 = nn.ReLU()
+ self.conv1 = nn.Conv2d(channels,
+ mid_channels,
+ kernel_size=1,
+ padding=0)
+ self.conv2 = nn.Conv2d(mid_channels,
+ 1,
+ kernel_size=3,
+ padding=1,
+ bias=False)
+ self.sigmoid = nn.Sigmoid()
+ nn.init.constant_(self.conv1.bias, 0)
+
+ def forward(self, x):
+ y = self.relu1(x)
+ y = self.conv1(y)
+ y = self.conv2(y)
+ y = self.sigmoid(y)
+
+ return x * y
+
+
+class CDCM(nn.Module):
+ """
+ Compact Dilation Convolution based Module
+ """
+ def __init__(self, in_channels, out_channels):
+ super().__init__()
+
+ self.relu1 = nn.ReLU()
+ self.conv1 = nn.Conv2d(in_channels,
+ out_channels,
+ kernel_size=1,
+ padding=0)
+ self.conv2_1 = nn.Conv2d(out_channels,
+ out_channels,
+ kernel_size=3,
+ dilation=5,
+ padding=5,
+ bias=False)
+ self.conv2_2 = nn.Conv2d(out_channels,
+ out_channels,
+ kernel_size=3,
+ dilation=7,
+ padding=7,
+ bias=False)
+ self.conv2_3 = nn.Conv2d(out_channels,
+ out_channels,
+ kernel_size=3,
+ dilation=9,
+ padding=9,
+ bias=False)
+ self.conv2_4 = nn.Conv2d(out_channels,
+ out_channels,
+ kernel_size=3,
+ dilation=11,
+ padding=11,
+ bias=False)
+ nn.init.constant_(self.conv1.bias, 0)
+
+ def forward(self, x):
+ x = self.relu1(x)
+ x = self.conv1(x)
+ x1 = self.conv2_1(x)
+ x2 = self.conv2_2(x)
+ x3 = self.conv2_3(x)
+ x4 = self.conv2_4(x)
+ return x1 + x2 + x3 + x4
+
+
+class MapReduce(nn.Module):
+ """
+ Reduce feature maps into a single edge map
+ """
+ def __init__(self, channels):
+ super().__init__()
+ self.conv = nn.Conv2d(channels, 1, kernel_size=1, padding=0)
+ nn.init.constant_(self.conv.bias, 0)
+
+ def forward(self, x):
+ return self.conv(x)
+
+
+class PDCBlock(nn.Module):
+ def __init__(self, pdc, inplane, ouplane, stride=1):
+ super().__init__()
+ self.stride = stride
+
+ self.stride = stride
+ if self.stride > 1:
+ self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
+ self.shortcut = nn.Conv2d(inplane,
+ ouplane,
+ kernel_size=1,
+ padding=0)
+ self.conv1 = Conv2d(pdc,
+ inplane,
+ inplane,
+ kernel_size=3,
+ padding=1,
+ groups=inplane,
+ bias=False)
+ self.relu2 = nn.ReLU()
+ self.conv2 = nn.Conv2d(inplane,
+ ouplane,
+ kernel_size=1,
+ padding=0,
+ bias=False)
+
+ def forward(self, x):
+ if self.stride > 1:
+ x = self.pool(x)
+ y = self.conv1(x)
+ y = self.relu2(y)
+ y = self.conv2(y)
+ if self.stride > 1:
+ x = self.shortcut(x)
+ y = y + x
+ return y
+
+
+class PDCBlock_converted(nn.Module):
+ """
+ CPDC, APDC can be converted to vanilla 3x3 convolution
+ RPDC can be converted to vanilla 5x5 convolution
+ """
+ def __init__(self, pdc, inplane, ouplane, stride=1):
+ super().__init__()
+ self.stride = stride
+
+ if self.stride > 1:
+ self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
+ self.shortcut = nn.Conv2d(inplane,
+ ouplane,
+ kernel_size=1,
+ padding=0)
+ if pdc == 'rd':
+ self.conv1 = nn.Conv2d(inplane,
+ inplane,
+ kernel_size=5,
+ padding=2,
+ groups=inplane,
+ bias=False)
+ else:
+ self.conv1 = nn.Conv2d(inplane,
+ inplane,
+ kernel_size=3,
+ padding=1,
+ groups=inplane,
+ bias=False)
+ self.relu2 = nn.ReLU()
+ self.conv2 = nn.Conv2d(inplane,
+ ouplane,
+ kernel_size=1,
+ padding=0,
+ bias=False)
+
+ def forward(self, x):
+ if self.stride > 1:
+ x = self.pool(x)
+ y = self.conv1(x)
+ y = self.relu2(y)
+ y = self.conv2(y)
+ if self.stride > 1:
+ x = self.shortcut(x)
+ y = y + x
+ return y
+
+
+class PiDiNet(nn.Module):
+ def __init__(self,
+ inplane,
+ pdcs,
+ dil=None,
+ sa=False,
+ convert=False,
+ mean=[0.485, 0.456, 0.406],
+ std=[0.229, 0.224, 0.225]):
+ super().__init__()
+ self.sa = sa
+ if dil is not None:
+ assert isinstance(dil, int), 'dil should be an int'
+ self.dil = dil
+ self.mean = mean
+ self.std = std
+
+ self.fuseplanes = []
+
+ self.inplane = inplane
+ if convert:
+ if pdcs[0] == 'rd':
+ init_kernel_size = 5
+ init_padding = 2
+ else:
+ init_kernel_size = 3
+ init_padding = 1
+ self.init_block = nn.Conv2d(3,
+ self.inplane,
+ kernel_size=init_kernel_size,
+ padding=init_padding,
+ bias=False)
+ block_class = PDCBlock_converted
+ else:
+ self.init_block = Conv2d(pdcs[0],
+ 3,
+ self.inplane,
+ kernel_size=3,
+ padding=1)
+ block_class = PDCBlock
+
+ self.block1_1 = block_class(pdcs[1], self.inplane, self.inplane)
+ self.block1_2 = block_class(pdcs[2], self.inplane, self.inplane)
+ self.block1_3 = block_class(pdcs[3], self.inplane, self.inplane)
+ self.fuseplanes.append(self.inplane) # C
+
+ inplane = self.inplane
+ self.inplane = self.inplane * 2
+ self.block2_1 = block_class(pdcs[4], inplane, self.inplane, stride=2)
+ self.block2_2 = block_class(pdcs[5], self.inplane, self.inplane)
+ self.block2_3 = block_class(pdcs[6], self.inplane, self.inplane)
+ self.block2_4 = block_class(pdcs[7], self.inplane, self.inplane)
+ self.fuseplanes.append(self.inplane) # 2C
+
+ inplane = self.inplane
+ self.inplane = self.inplane * 2
+ self.block3_1 = block_class(pdcs[8], inplane, self.inplane, stride=2)
+ self.block3_2 = block_class(pdcs[9], self.inplane, self.inplane)
+ self.block3_3 = block_class(pdcs[10], self.inplane, self.inplane)
+ self.block3_4 = block_class(pdcs[11], self.inplane, self.inplane)
+ self.fuseplanes.append(self.inplane) # 4C
+
+ self.block4_1 = block_class(pdcs[12],
+ self.inplane,
+ self.inplane,
+ stride=2)
+ self.block4_2 = block_class(pdcs[13], self.inplane, self.inplane)
+ self.block4_3 = block_class(pdcs[14], self.inplane, self.inplane)
+ self.block4_4 = block_class(pdcs[15], self.inplane, self.inplane)
+ self.fuseplanes.append(self.inplane) # 4C
+
+ self.conv_reduces = nn.ModuleList()
+ if self.sa and self.dil is not None:
+ self.attentions = nn.ModuleList()
+ self.dilations = nn.ModuleList()
+ for i in range(4):
+ self.dilations.append(CDCM(self.fuseplanes[i], self.dil))
+ self.attentions.append(CSAM(self.dil))
+ self.conv_reduces.append(MapReduce(self.dil))
+ elif self.sa:
+ self.attentions = nn.ModuleList()
+ for i in range(4):
+ self.attentions.append(CSAM(self.fuseplanes[i]))
+ self.conv_reduces.append(MapReduce(self.fuseplanes[i]))
+ elif self.dil is not None:
+ self.dilations = nn.ModuleList()
+ for i in range(4):
+ self.dilations.append(CDCM(self.fuseplanes[i], self.dil))
+ self.conv_reduces.append(MapReduce(self.dil))
+ else:
+ for i in range(4):
+ self.conv_reduces.append(MapReduce(self.fuseplanes[i]))
+
+ self.classifier = nn.Conv2d(4, 1, kernel_size=1) # has bias
+ nn.init.constant_(self.classifier.weight, 0.25)
+ nn.init.constant_(self.classifier.bias, 0)
+
+ def get_weights(self):
+ conv_weights = []
+ bn_weights = []
+ relu_weights = []
+ for pname, p in self.named_parameters():
+ if 'bn' in pname:
+ bn_weights.append(p)
+ elif 'relu' in pname:
+ relu_weights.append(p)
+ else:
+ conv_weights.append(p)
+
+ return conv_weights, bn_weights, relu_weights
+
+ def forward(self, x):
+ """x: [B, 3, H, W] within range [0, 1].
+ """
+ x = (x - x.new_tensor(self.mean).view(1, -1, 1, 1)) / \
+ x.new_tensor(self.std).view(1, -1, 1, 1)
+ h, w = x.size()[2:]
+
+ x = self.init_block(x)
+
+ x1 = self.block1_1(x)
+ x1 = self.block1_2(x1)
+ x1 = self.block1_3(x1)
+
+ x2 = self.block2_1(x1)
+ x2 = self.block2_2(x2)
+ x2 = self.block2_3(x2)
+ x2 = self.block2_4(x2)
+
+ x3 = self.block3_1(x2)
+ x3 = self.block3_2(x3)
+ x3 = self.block3_3(x3)
+ x3 = self.block3_4(x3)
+
+ x4 = self.block4_1(x3)
+ x4 = self.block4_2(x4)
+ x4 = self.block4_3(x4)
+ x4 = self.block4_4(x4)
+
+ x_fuses = []
+ if self.sa and self.dil is not None:
+ for i, xi in enumerate([x1, x2, x3, x4]):
+ x_fuses.append(self.attentions[i](self.dilations[i](xi)))
+ elif self.sa:
+ for i, xi in enumerate([x1, x2, x3, x4]):
+ x_fuses.append(self.attentions[i](xi))
+ elif self.dil is not None:
+ for i, xi in enumerate([x1, x2, x3, x4]):
+ x_fuses.append(self.dilations[i](xi))
+ else:
+ x_fuses = [x1, x2, x3, x4]
+
+ e1 = self.conv_reduces[0](x_fuses[0])
+ e1 = F.interpolate(e1, (h, w), mode='bilinear', align_corners=False)
+
+ e2 = self.conv_reduces[1](x_fuses[1])
+ e2 = F.interpolate(e2, (h, w), mode='bilinear', align_corners=False)
+
+ e3 = self.conv_reduces[2](x_fuses[2])
+ e3 = F.interpolate(e3, (h, w), mode='bilinear', align_corners=False)
+
+ e4 = self.conv_reduces[3](x_fuses[3])
+ e4 = F.interpolate(e4, (h, w), mode='bilinear', align_corners=False)
+
+ outputs = [e1, e2, e3, e4]
+ output = self.classifier(torch.cat(outputs, dim=1))
+
+ outputs.append(output)
+ outputs = [torch.sigmoid(r) for r in outputs]
+ return outputs[-1]
+
+
+@ANNOTATORS.register_class()
+class PiDiAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ pretrained_model = cfg.get('PRETRAINED_MODEL', None)
+ vanilla_cnn = cfg.get('VANILLA_CNN', True)
+ pdcs = config_model_converted(
+ 'carv4') if vanilla_cnn else config_model('carv4')
+ self.model = PiDiNet(60, pdcs, dil=24, sa=True,
+ convert=vanilla_cnn).eval()
+ if pretrained_model:
+ with FS.get_from(pretrained_model, wait_finish=True) as local_path:
+ state = torch.load(local_path,
+ map_location='cpu')['state_dict']
+ if vanilla_cnn:
+ state = convert_pidinet(state, 'carv4')
+ state = {
+ k[len('module.'):] if k.startswith('module.') else k: v
+ for k, v in state.items()
+ }
+ self.model.load_state_dict(state)
+
+ @torch.no_grad()
+ @torch.inference_mode()
+ @torch.autocast('cuda', enabled=False)
+ def forward(self, image, return_grayscale=False):
+ is_batch = False if len(image.shape) == 3 else True
+ if isinstance(image, torch.Tensor):
+ if len(image.shape) == 3:
+ image = rearrange(image, 'h w c -> 1 c h w')
+ elif len(image.shape) == 4:
+ image = rearrange(image, 'b h w c -> b c h w')
+ else:
+ raise "Unsurpport input image's shape"
+ elif isinstance(image, np.ndarray):
+ image = torch.from_numpy(image.copy()).float()
+ if len(image.shape) == 3:
+ image = rearrange(image, 'h w c -> 1 c h w')
+ elif len(image.shape) == 4:
+ image = rearrange(image, 'b h w c -> b c h w')
+ else:
+ raise "Unsurpport input image's shape"
+ else:
+ raise "Unsurpport input image's type"
+ image = image.float().div(255)
+ image = image.to(we.device_id)
+ edge = self.model(image)
+ edge = edge.squeeze(dim=1)
+ edge = 255 - (edge * 255.0).clip(0, 255) # return white background
+ edge = edge.cpu().numpy()
+ edge = edge.astype(np.uint8)
+ if not is_batch:
+ edge = edge.squeeze()
+ if not return_grayscale:
+ edge = edge[..., None].repeat(3, -1)
+ return edge
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ PiDiAnnotator.para_dict,
+ set_name=True)
diff --git a/scepter/modules/annotator/segmentation.py b/scepter/modules/annotator/segmentation.py
new file mode 100644
index 0000000..ddb2ae5
--- /dev/null
+++ b/scepter/modules/annotator/segmentation.py
@@ -0,0 +1,382 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import math
+import random
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+import torchvision.transforms as T
+from PIL import Image
+from scipy import ndimage
+from pycocotools import mask as mask_utils
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+from sklearn.cluster import KMeans
+from torchvision.ops.boxes import batched_nms
+
+
+def find_dominant_color(image, k=1):
+ pixels = image.reshape((-1, 3))
+ mask = (pixels != [0, 0, 0]).all(axis=1)
+ pixels = pixels[mask]
+ try:
+ kmeans = KMeans(n_clusters=k, n_init='auto')
+ kmeans.fit(pixels)
+ dominant_color = kmeans.cluster_centers_.astype(int)[0]
+ except:
+ dominant_color = np.array([255, 255, 255])
+ return dominant_color
+
+
+def cv2_resize_crop(image, resize_size, crop_size):
+ resize_height, resize_width = resize_size
+ crop_height, crop_width = crop_size
+
+ resized_image = cv2.resize(image, (resize_width, resize_height))
+
+ center_x, center_y = resize_width // 2, resize_height // 2
+ crop_start_x = max(center_x - crop_width // 2, 0)
+ crop_start_y = max(center_y - crop_height // 2, 0)
+ crop_end_x = crop_start_x + crop_width
+ crop_end_y = crop_start_y + crop_height
+
+ crop_end_x = min(crop_end_x, resize_width)
+ crop_end_y = min(crop_end_y, resize_height)
+
+ center_cropped_image = resized_image[crop_start_y:crop_end_y,
+ crop_start_x:crop_end_x]
+
+ return center_cropped_image
+
+
+@ANNOTATORS.register_class()
+class ESAMAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ try:
+ from efficient_sam.efficient_sam import build_efficient_sam
+ from segment_anything.utils.amg import (
+ batched_mask_to_box,
+ calculate_stability_score,
+ mask_to_rle_pytorch,
+ remove_small_regions,
+ rle_to_mask,
+ )
+ except:
+ raise NotImplementedError(
+ f'Please install efficient_sam and segment_anything modules.')
+
+ pretrained_model = cfg.get('PRETRAINED_MODEL', None)
+ if pretrained_model:
+ with FS.get_from(pretrained_model, wait_finish=True) as local_path:
+ self.efficient_sam_module = build_efficient_sam(
+ encoder_patch_embed_dim=384,
+ encoder_num_heads=6,
+ checkpoint=local_path).eval().to(we.device_id)
+ self.GRID_SIZE = cfg.get('GRID_SIZE', 16)
+ self.save_mode = cfg.get('SAVE_MODE', 'P')
+ self.use_dominant_color = cfg.get('USE_DOMINANT_COLOR', False)
+ self.return_mask = cfg.get('RETURN_MASK', False)
+
+ @torch.no_grad()
+ def get_predictions_given_embeddings_and_queries(self, img, points,
+ point_labels, model):
+ from segment_anything.utils.amg import calculate_stability_score
+ predicted_masks, predicted_iou = [], []
+ bs = 128
+ num = int(float(self.GRID_SIZE * self.GRID_SIZE) /
+ bs) if self.GRID_SIZE * self.GRID_SIZE % bs == 0 else int(
+ float(self.GRID_SIZE * self.GRID_SIZE) / bs) + 1
+ for i in range(num):
+ predicted_mask_item, predicted_iou_item = model(
+ img[None, ...], points[:, i * bs:(i + 1) * bs, ...],
+ point_labels[:, i * bs:(i + 1) * bs, :])
+ predicted_masks.append(predicted_mask_item)
+ predicted_iou.append(predicted_iou_item)
+ torch.cuda.empty_cache()
+ # # predicted: torch.Size([1, 1024, 3]) torch.Size([1, 1024, 3, 512, 512])
+ # print('predicted: ', predicted_iou.size(), predicted_masks.size())
+ predicted_masks = torch.cat(predicted_masks, dim=1)
+ predicted_iou = torch.cat(predicted_iou, dim=1)
+ sorted_ids = torch.argsort(predicted_iou, dim=-1, descending=True)
+ predicted_iou_scores = torch.take_along_dim(predicted_iou,
+ sorted_ids,
+ dim=2)
+ predicted_masks = torch.take_along_dim(predicted_masks,
+ sorted_ids[..., None, None],
+ dim=2)
+ predicted_masks = predicted_masks[0]
+ iou = predicted_iou_scores[0, :, 0]
+ index_iou = iou > 0.7
+ iou_ = iou[index_iou]
+ masks = predicted_masks[index_iou]
+ score = calculate_stability_score(masks, 0.0, 1.0)
+ score = score[:, 0]
+ index = score > 0.9
+ masks = masks[index]
+ iou_ = iou_[index]
+ masks = torch.ge(masks, 0.0)
+ return masks, iou_
+
+ def singel_mask_to_rle(self, mask):
+ rle = mask_utils.encode(
+ np.array(mask[:, :, None], order='F', dtype='uint8'))[0]
+ rle['counts'] = rle['counts'].decode('utf-8')
+ return rle
+
+ def process_small_region(self, rles):
+ from segment_anything.utils.amg import rle_to_mask, remove_small_regions, \
+ batched_mask_to_box, mask_to_rle_pytorch
+ new_masks = []
+ scores = []
+ min_area = 100
+ nms_thresh = 0.7
+ for rle in rles:
+ mask = rle_to_mask(rle[0])
+
+ mask, changed = remove_small_regions(mask, min_area, mode='holes')
+ unchanged = not changed
+ mask, changed = remove_small_regions(mask,
+ min_area,
+ mode='islands')
+ unchanged = unchanged and not changed
+
+ new_masks.append(torch.as_tensor(mask).unsqueeze(0))
+ # Give score=0 to changed masks and score=1 to unchanged masks
+ # so NMS will prefer ones that didn't need postprocessing
+ scores.append(float(unchanged))
+
+ # Recalculate boxes and remove any new duplicates
+ masks = torch.cat(new_masks, dim=0).to(we.device_id)
+ boxes = batched_mask_to_box(masks)
+ keep_by_nms = batched_nms(
+ boxes.float(),
+ torch.as_tensor(scores).to(we.device_id),
+ torch.zeros_like(boxes[:, 0]), # categories
+ iou_threshold=nms_thresh,
+ )
+
+ # Only recalculate RLEs for masks that have changed
+ for i_mask in keep_by_nms:
+ if scores[i_mask] == 0.0:
+ mask_torch = masks[i_mask].unsqueeze(0)
+ rles[i_mask] = mask_to_rle_pytorch(mask_torch)
+ masks = [rle_to_mask(rles[i][0]) for i in keep_by_nms]
+ return masks
+
+ def run_everything_ours(self, img_tensor, model):
+ from segment_anything.utils.amg import mask_to_rle_pytorch
+ img_tensor = img_tensor.squeeze(0)
+ _, original_image_h, original_image_w = img_tensor.shape
+ xy = []
+ for i in range(self.GRID_SIZE):
+ curr_x = 0.5 + i / self.GRID_SIZE * original_image_w
+ for j in range(self.GRID_SIZE):
+ curr_y = 0.5 + j / self.GRID_SIZE * original_image_h
+ xy.append([curr_x, curr_y])
+
+ xy = torch.from_numpy(np.array(xy))
+ points = xy
+ num_pts = xy.shape[0]
+ point_labels = torch.ones(num_pts, 1)
+ with torch.no_grad():
+ predicted_masks, predicted_iou = self.get_predictions_given_embeddings_and_queries(
+ img_tensor,
+ points.reshape(1, num_pts, 1, 2).to(we.device_id),
+ point_labels.reshape(1, num_pts, 1).to(we.device_id),
+ model,
+ )
+ # print('predicted_masks: ', predicted_masks[0][0:1].dtype, predicted_masks[0][0:1].device)
+ rle = [mask_to_rle_pytorch(m[0:1]) for m in predicted_masks]
+ # transform to numpy
+ size, counts = [], []
+ for rle_item in rle:
+ size.append(rle_item[0]['size'])
+ counts += rle_item[0]['counts']
+ counts += '#'
+ predicted_masks = self.process_small_region(rle)
+ return predicted_masks
+
+ def forward(self, image, return_mask=None):
+ return_mask = return_mask if return_mask is not None else self.return_mask
+ if isinstance(image, Image.Image):
+ image = np.array(image)
+ elif isinstance(image, torch.Tensor):
+ image = image.detach().cpu().numpy()
+ elif isinstance(image, np.ndarray):
+ image = image.copy()
+ else:
+ raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+ h, w = image.shape[:2]
+ max_rate = max(float(w) / 1024.0, float(h) / 1024.0)
+ w_ori = int(float(w) / max_rate)
+ h_ori = int(float(h) / max_rate)
+ # image = T.ToTensor()(T.Resize((h_ori, w_ori))(Image.fromarray(image)))
+ # image_pad = T.Pad((0, 0, 1024 - w_ori, 1024 - h_ori))(image)
+ image_pad = T.Pad((0, 0, 1024 - w_ori, 1024 - h_ori))(T.Resize(
+ (h_ori, w_ori))(Image.fromarray(image)))
+ input_image = T.ToTensor()(image_pad)
+ input_image = input_image.unsqueeze(0).to(we.device_id)
+
+ mask_efficient_sam_vits = self.run_everything_ours(
+ input_image, self.efficient_sam_module)
+ annos = []
+ mask_efficient_sam_vits = sorted(list(mask_efficient_sam_vits),
+ key=lambda m: int(m.sum()),
+ reverse=True)
+ mask_efficient_sam_vits = mask_efficient_sam_vits[:256]
+ for mask in mask_efficient_sam_vits:
+ mask_item = mask_utils.encode(
+ np.array(mask[:, :, None], order='F', dtype='uint8'))[0]
+ mask_item['counts'] = mask_item['counts'].decode('utf-8')
+ mask_area = int(mask.sum())
+ annos.append({'mask': mask_item, 'mask_area': mask_area})
+
+ annos = sorted(annos, key=lambda x: x['mask_area'], reverse=True)
+ seg_img = None
+ dominant_palette = []
+ image_pad_np = np.array(image_pad)
+ for idx, anno in enumerate(annos):
+ color = idx
+ if idx > 255:
+ break
+ mask = np.array(mask_utils.decode(anno['mask'])).astype(np.uint8)
+ h, w = mask.shape[:2]
+ if seg_img is None:
+ seg_img = np.ones((h, w, 3)) * 255
+ if self.use_dominant_color:
+ masked_image = cv2.bitwise_and(image_pad_np,
+ image_pad_np,
+ mask=mask)
+ dominant_color = find_dominant_color(masked_image).tolist()
+ dominant_palette.append(dominant_color)
+ seg_img[mask.astype(bool)] = [color, color, color]
+ seg_img = Image.fromarray(seg_img.astype(np.uint8)).convert('L')
+
+ resize_rate = max(float(h_ori) / 1024.0, float(w_ori) / 1024.0)
+ h_new = int(float(h_ori) / resize_rate)
+ w_new = int(float(w_ori) / resize_rate)
+ seg_img = seg_img.crop((0, 0, w_new, h_new))
+ if self.save_mode == 'P':
+ palette = []
+ for i in range(256):
+ if not self.use_dominant_color:
+ palette_item = [random.randint(0, 255) for _ in range(3)]
+ else:
+ palette_item = dominant_palette[i] if i < len(
+ dominant_palette) else [255, 255, 255]
+ palette += palette_item
+ seg_img = seg_img.convert('P')
+ seg_img.putpalette(palette)
+ seg_rgb_img = seg_img.convert('RGB')
+ if return_mask:
+ return {
+ 'image': np.array(seg_rgb_img),
+ 'mask': np.array(seg_img)
+ }
+ else:
+ return np.array(seg_rgb_img)
+ else:
+ return np.array(seg_img)
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ ESAMAnnotator.para_dict,
+ set_name=True)
+
+
+
+
+@ANNOTATORS.register_class()
+class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ from segment_anything import sam_model_registry, SamPredictor
+ from segment_anything.utils.transforms import ResizeLongestSide
+
+ self.transform = ResizeLongestSide(1024)
+ self.task_type = cfg.get('TASK_TYPE', 'input_box')
+ self.sam_model = cfg.get('SAM_MODEL', 'vit_b')
+ pretrained_model = cfg.get('PRETRAINED_MODEL', 'sam_vit_b_01ec64.pth')
+
+ if pretrained_model:
+ with FS.get_from(pretrained_model, wait_finish=True) as local_path:
+ seg_model = sam_model_registry[self.sam_model](checkpoint=local_path).eval().to(we.device_id)
+ self.sam_predictor = SamPredictor(seg_model)
+
+ def forward(self, image, input_box=None, mask=None, task_type=None, multimask_output=False):
+ task_type = task_type if task_type is not None else self.task_type
+
+ if isinstance(image, Image.Image):
+ image = np.array(image)
+ elif isinstance(image, torch.Tensor):
+ image = image.detach().cpu().numpy()
+ elif isinstance(image, np.ndarray):
+ image = image.copy()
+ else:
+ raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+
+ if mask is not None:
+ if isinstance(mask, Image.Image):
+ mask = np.array(mask)
+ elif isinstance(mask, torch.Tensor):
+ mask = mask.detach().cpu().numpy()
+ elif isinstance(mask, np.ndarray):
+ mask = mask.copy()
+ else:
+ raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
+
+ original_size = image.shape[:2]
+ if task_type == 'mask_point':
+ scribble = mask.transpose(2, 1, 0)[0]
+ labeled_array, num_features = ndimage.label(scribble >= 255)
+ centers = ndimage.center_of_mass(scribble, labeled_array, range(1, num_features + 1))
+ point_coords = np.array(centers)
+ point_labels = np.array([1] * len(centers))
+ sample = {'point_coords': point_coords, 'point_labels': point_labels}
+
+ elif task_type == 'mask_box':
+ scribble = mask.transpose(2, 1, 0)[0]
+ labeled_array, num_features = ndimage.label(scribble >= 255)
+ centers = ndimage.center_of_mass(scribble, labeled_array, range(1, num_features + 1))
+ centers = np.array(centers)
+ ### (x1, y1, x2, y2)
+ x_min = centers[:, 0].min()
+ x_max = centers[:, 0].max()
+ y_min = centers[:, 1].min()
+ y_max = centers[:, 1].max()
+ bbox = np.array([x_min, y_min, x_max, y_max])
+ sample = {'box': bbox}
+
+ elif task_type == 'input_box':
+ if isinstance(input_box, list):
+ input_box = np.array(input_box)
+ sample = {'box': input_box}
+
+ self.sam_predictor.set_image(image)
+ masks, scores, logits = self.sam_predictor.predict(**sample, multimask_output=True)
+ index = np.argmax(scores)
+
+ ret_data = {
+ "mask": (masks[index]* 255).astype(np.uint8),
+ "score": scores[index]
+ }
+ return ret_data
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ SAMAnnotatorDraw.para_dict,
+ set_name=True)
\ No newline at end of file
diff --git a/scepter/modules/annotator/sketch.py b/scepter/modules/annotator/sketch.py
new file mode 100644
index 0000000..375d5d0
--- /dev/null
+++ b/scepter/modules/annotator/sketch.py
@@ -0,0 +1,160 @@
+# -*- coding: utf-8 -*-
+
+import math
+from abc import ABCMeta
+
+import cv2
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torchvision.transforms as TT
+from einops import rearrange
+from scepter.modules.annotator.base_annotator import BaseAnnotator
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+
+
+class SketchNet(nn.Module):
+ def __init__(self, mean, std):
+ assert isinstance(mean, float) and isinstance(std, float)
+ super().__init__()
+ self.mean = mean
+ self.std = std
+
+ # layers
+ self.layers = nn.Sequential(nn.Conv2d(1, 48, 5, 2, 2),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(48, 128, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(128, 128, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(128, 128, 3, 2, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(128, 256, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(256, 256, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(256, 256, 3, 2, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(256, 512, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(512, 1024, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(1024, 1024, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(1024, 1024, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(1024, 1024, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(1024, 512, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(512, 256, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.ConvTranspose2d(256, 256, 4, 2, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(256, 256, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(256, 128, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.ConvTranspose2d(128, 128, 4, 2, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(128, 128, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(128, 48, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.ConvTranspose2d(48, 48, 4, 2, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(48, 24, 3, 1, 1),
+ nn.ReLU(inplace=True),
+ nn.Conv2d(24, 1, 3, 1, 1), nn.Sigmoid())
+
+ def forward(self, x):
+ """x: [B, 1, H, W] within range [0, 1]. Sketch pixels in dark color.
+ """
+ x = (x - self.mean) / self.std
+ return self.layers(x)
+
+
+@ANNOTATORS.register_class()
+class SketchAnnotator(BaseAnnotator, metaclass=ABCMeta):
+ para_dict = {}
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ pretrained_model = cfg.get('PRETRAINED_MODEL', None)
+ self.model = SketchNet(mean=0.9664114577640158,
+ std=0.0858381272736797).eval()
+ if pretrained_model:
+ with FS.get_from(pretrained_model, wait_finish=True) as local_path:
+ state = torch.load(local_path, map_location='cpu')
+ self.model.load_state_dict(state)
+
+ @torch.no_grad()
+ @torch.inference_mode()
+ @torch.autocast('cuda', enabled=False)
+ def forward(self, image):
+ is_batch = False if len(image.shape) == 3 else True
+ if isinstance(image, torch.Tensor):
+ if len(image.shape) == 3:
+ if torch.equal(image[:, :, 0], image[:, :, 1]) and torch.equal(
+ image[:, :, 1], image[:, :, 2]):
+ image = image[:, :, 0]
+ else:
+ raise "Unsurpport input image's shape and each channel is different"
+ elif len(image.shape) == 4:
+ if (torch.equal(image[:, :, :, 0], image[:, :, :, 1])
+ and torch.equal(image[:, :, :, 1], image[:, :, :, 2])):
+ image = image[:, :, :, 0]
+ else:
+ raise "Unsurpport input image's shape and each channel is different"
+ if len(image.shape) == 2:
+ image = rearrange(image, 'h w -> 1 h w')
+ B, H, W = image.shape
+ elif len(image.shape) == 3:
+ B, H, W = image.shape
+ else:
+ raise "Unsurpport input image's shape"
+ elif isinstance(image, np.ndarray):
+ image = image.copy()
+ if len(image.shape) == 3:
+ if np.array_equal(image[:, :, 0],
+ image[:, :, 1]) and np.array_equal(
+ image[:, :, 1], image[:, :, 2]):
+ image = image[:, :, 0]
+ else:
+ raise "Unsurpport input image's shape and each channel is different"
+ elif len(image.shape) == 4:
+ if (np.array_equal(image[:, :, :, 0], image[:, :, :, 1]) and
+ np.array_equal(image[:, :, :, 1], image[:, :, :, 2])):
+ image = image[:, :, :, 0]
+ else:
+ raise "Unsurpport input image's shape and each channel is different"
+ image = torch.from_numpy(image).float()
+ if len(image.shape) == 2:
+ image = rearrange(image, 'h w -> 1 1 h w')
+ elif len(image.shape) == 3:
+ image = rearrange(image, 'b h w -> b 1 h w')
+ else:
+ raise "Unsurpport input image's shape"
+ else:
+ raise "Unsurpport input image's type"
+ image = image.float().div(255)
+ image = image.to(we.device_id)
+ edge = self.model(image)
+ edge = edge.squeeze(dim=1)
+ edge = (edge * 255.0).clip(0, 255)
+ edge = edge.cpu().numpy()
+ edge = edge.astype(np.uint8)
+ if not is_batch:
+ edge = edge.squeeze()
+ return edge[..., None].repeat(3, -1)
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('ANNOTATORS',
+ __class__.__name__,
+ SketchAnnotator.para_dict,
+ set_name=True)
diff --git a/scepter/modules/data/dataset/base_dataset.py b/scepter/modules/data/dataset/base_dataset.py
index e985201..e35c805 100644
--- a/scepter/modules/data/dataset/base_dataset.py
+++ b/scepter/modules/data/dataset/base_dataset.py
@@ -82,6 +82,8 @@ class BaseDataset(Dataset, metaclass=ABCMeta):
overwrite=False)
self.worker_id = worker_id
self.logger = self.worker_logger
+ self.local_we["seed"] += (worker_id + we.rank)
+ self.seed = self.local_we["seed"]
we.set_env(self.local_we)
@abstractmethod
diff --git a/scepter/modules/data/dataset/registry.py b/scepter/modules/data/dataset/registry.py
index 573c0a2..3906d41 100644
--- a/scepter/modules/data/dataset/registry.py
+++ b/scepter/modules/data/dataset/registry.py
@@ -258,6 +258,8 @@ class DataObject(object):
if sampler_name == 'MixtureOfSamplers':
subsampler_configs = self.data_sampler_config.get(
'SUB_SAMPLERS', [])
+ keep_order = self.data_sampler_config.get(
+ 'KEEP_ORDER', False)
subsamplers = list()
subsampler_probs = list()
for ssconfig in subsampler_configs:
@@ -277,7 +279,7 @@ class DataObject(object):
subsampler_probs.append(prob)
self.batch_sampler = MixtureOfSamplers(subsamplers,
subsampler_probs, rank,
- seed)
+ seed, keep_order = keep_order)
elif sampler_name == 'MultiLevelBatchSampler':
self.batch_sampler = self._instantiate_multi_level_batch_sampler(
self.data_sampler_config, self.batch_size, rank, seed)
diff --git a/scepter/modules/data/sampler/sampler.py b/scepter/modules/data/sampler/sampler.py
index 50ea158..e6bb7e1 100644
--- a/scepter/modules/data/sampler/sampler.py
+++ b/scepter/modules/data/sampler/sampler.py
@@ -523,12 +523,15 @@ class MultiLevelBatchSampler(BaseSampler):
class MixtureOfSamplers(BaseSampler):
para_dict = {'SUB_SAMPLERS': []}
- def __init__(self, samplers, probabilities, rank=0, seed=8888):
+ def __init__(self, samplers, probabilities, rank=0, seed=8888, keep_order = False):
self.samplers = samplers
self.iterators = [iter(u) for u in samplers]
self.probabilities = probabilities
self.seed = seed
- self.rng = np.random.default_rng(seed + rank)
+ if keep_order:
+ self.rng = np.random.default_rng(seed)
+ else:
+ self.rng = np.random.default_rng(seed + rank)
def __iter__(self):
while True:
diff --git a/scepter/modules/inference/control_inference.py b/scepter/modules/inference/control_inference.py
index 1d08e14..21cadb8 100644
--- a/scepter/modules/inference/control_inference.py
+++ b/scepter/modules/inference/control_inference.py
@@ -13,12 +13,6 @@ from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
-try:
- from swift import SwiftModel
-except Exception:
- warnings.warn('Import swift failed, please check it.')
-
-
class ControlInference():
def __init__(self, logger=None):
self.logger = logger
@@ -26,6 +20,11 @@ class ControlInference():
# @classmethod
def unregister_controllers(self, control_model_ins, diffusion_model):
+ try:
+ from swift import SwiftModel
+ except Exception:
+ warnings.warn('Import swift failed, please check it.')
+
self.logger.info('Unloading control model')
if isinstance(diffusion_model['model'], SwiftModel):
if (hasattr(diffusion_model['model'].base_model, 'control_blocks')
@@ -42,6 +41,11 @@ class ControlInference():
# @classmethod
def register_controllers(self, control_model_ins, diffusion_model):
+ try:
+ from swift import SwiftModel
+ except Exception:
+ warnings.warn('Import swift failed, please check it.')
+
self.logger.info('Loading control model')
if control_model_ins is None or control_model_ins == '':
self.unregister_controllers(control_model_ins, diffusion_model)
diff --git a/scepter/modules/inference/diffusion_inference.py b/scepter/modules/inference/diffusion_inference.py
index ff1a427..f00989a 100644
--- a/scepter/modules/inference/diffusion_inference.py
+++ b/scepter/modules/inference/diffusion_inference.py
@@ -84,14 +84,13 @@ class DiffusionInference():
def redefine_paras(self, cfg):
if cfg.get('PRETRAINED_MODEL', None):
- assert FS.isfile(cfg.PRETRAINED_MODEL)
with FS.get_from(cfg.PRETRAINED_MODEL,
wait_finish=True) as local_path:
if local_path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_path)
else:
- sd = torch.load(local_path, map_location='cpu')
+ sd = torch.load(local_path, map_location='cpu', weights_only=True)
first_stage_model_path = os.path.join(
os.path.dirname(local_path), 'first_stage_model.pth')
cond_stage_model_path = os.path.join(
@@ -311,7 +310,7 @@ class DiffusionInference():
module_paras = {}
if cfg is not None:
self.paras = cfg.PARAS
- self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
+ self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict)) else v for k, v in cfg.INPUT.items()}
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
module_paras = cfg.MODULES_PARAS
return module_paras
diff --git a/scepter/modules/inference/flux_inference.py b/scepter/modules/inference/flux_inference.py
new file mode 100644
index 0000000..4aefc43
--- /dev/null
+++ b/scepter/modules/inference/flux_inference.py
@@ -0,0 +1,216 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import copy
+import math
+import random
+import torch
+from scepter.modules.utils.distribute import we
+from .control_inference import ControlInference
+from .diffusion_inference import DiffusionInference, get_model
+from .tuner_inference import TunerInference
+from scepter.modules.model.registry import DIFFUSIONS, TOKENIZERS
+
+
+class FluxInference(DiffusionInference):
+ def __init__(self, logger=None):
+ self.logger = logger
+ self.is_redefine_paras = False
+ self.loaded_model = {}
+ self.loaded_model_name = [
+ 'diffusion_model', 'first_stage_model', 'cond_stage_model'
+ ]
+ self.tuner_infer = TunerInference(self.logger)
+ self.control_infer = ControlInference(self.logger)
+
+ def init_from_cfg(self, cfg):
+ self.name = cfg.NAME
+ self.is_default = cfg.get('IS_DEFAULT', False)
+ module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
+ assert cfg.have('MODEL')
+ if self.is_redefine_paras:
+ cfg.MODEL = self.redefine_paras(cfg.MODEL)
+ self.diffusion_model = self.infer_model(
+ cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
+ 'DIFFUSION_MODEL',
+ None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
+ self.first_stage_model = self.infer_model(
+ cfg.MODEL.FIRST_STAGE_MODEL,
+ module_paras.get(
+ 'FIRST_STAGE_MODEL',
+ None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
+ self.cond_stage_model = self.infer_model(
+ cfg.MODEL.COND_STAGE_MODEL,
+ module_paras.get(
+ 'COND_STAGE_MODEL',
+ None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
+ self.refiner_cond_model = self.infer_model(
+ cfg.MODEL.REFINER_COND_MODEL,
+ module_paras.get(
+ 'REFINER_COND_MODEL',
+ None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
+ self.refiner_diffusion_model = self.infer_model(
+ cfg.MODEL.REFINER_MODEL, module_paras.get(
+ 'REFINER_MODEL',
+ None)) if cfg.MODEL.have('REFINER_MODEL') else None
+ self.tokenizer = TOKENIZERS.build(
+ cfg.MODEL.TOKENIZER,
+ logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
+
+ if self.tokenizer is not None:
+ self.cond_stage_model['cfg'].KWARGS = {
+ 'vocab_size': self.tokenizer.vocab_size
+ }
+ self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION, logger=self.logger) \
+ if cfg.MODEL.have('DIFFUSION') else None
+ assert self.diffusion is not None
+
+ @torch.no_grad()
+ def encode_first_stage(self, x, **kwargs):
+ _, dtype = self.get_function_info(self.first_stage_model, 'encode')
+ with torch.autocast('cuda',
+ enabled= dtype in ('float16', 'bfloat16'),
+ dtype=getattr(torch, dtype)):
+ z = get_model(self.first_stage_model).encode(x)
+ if isinstance(z, (tuple, list)):
+ z = z[0]
+ return z
+
+ @torch.no_grad()
+ def decode_first_stage(self, z):
+ _, dtype = self.get_function_info(self.first_stage_model, 'decode')
+ with torch.autocast('cuda',
+ enabled=dtype in ('float16', 'bfloat16'),
+ dtype=getattr(torch, dtype)):
+ return get_model(self.first_stage_model).decode(z)
+
+ @torch.no_grad()
+ def __call__(self,
+ input,
+ num_samples=1,
+ cat_uc=True,
+ tuner_model=None,
+ control_model=None,
+ **kwargs):
+
+ value_input = copy.deepcopy(self.input)
+ value_input.update(input)
+ print(value_input)
+ height, width = value_input['target_size_as_tuple']
+ value_output = copy.deepcopy(self.output)
+ # register tuner
+ if tuner_model is not None and tuner_model != '' and len(
+ tuner_model) > 0:
+ if not isinstance(tuner_model, list):
+ tuner_model = [tuner_model]
+ self.dynamic_load(self.diffusion_model, 'diffusion_model')
+ self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
+ self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
+ self.cond_stage_model)
+ self.dynamic_unload(self.diffusion_model,
+ 'diffusion_model',
+ skip_loaded=True)
+ self.dynamic_unload(self.cond_stage_model,
+ 'cond_stage_model',
+ skip_loaded=True)
+
+ # cond stage
+ self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
+ function_name, dtype = self.get_function_info(self.cond_stage_model)
+ with torch.autocast('cuda',
+ enabled=dtype == 'float16',
+ dtype=getattr(torch, dtype)):
+ ctx = getattr(get_model(self.cond_stage_model),
+ function_name)(value_input['prompt'])
+ self.dynamic_unload(self.cond_stage_model,
+ 'cond_stage_model',
+ skip_loaded=True)
+
+ # get noise
+ seed = kwargs.pop('seed', -1)
+ g = torch.Generator(device=we.device_id)
+ seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
+ g.manual_seed(seed)
+ if 'seed' in value_output:
+ value_output['seed'] = seed
+ for sample_id in range(num_samples):
+ if self.diffusion_model is not None:
+ noise = torch.randn(
+ num_samples,
+ 16,
+ # allow for packing
+ 2 * math.ceil(height / 16),
+ 2 * math.ceil(width / 16),
+ device=we.device_id,
+ dtype=getattr(torch, dtype),
+ generator=torch.Generator(device=we.device_id).manual_seed(seed),
+ )
+ self.dynamic_load(self.diffusion_model, 'diffusion_model')
+ # UNet use input n_prompt
+ function_name, dtype = self.get_function_info(
+ self.diffusion_model)
+ with torch.autocast('cuda',
+ enabled= dtype in ('float16', 'bfloat16'),
+ dtype=getattr(torch, dtype)):
+ solver_sample = value_input.get('sample', 'flow_eluer')
+ sample_steps = value_input.get('sample_steps', 20)
+ guide_scale = value_input.get('guide_scale', 3.5)
+ if guide_scale is not None:
+ guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device,
+ dtype=noise.dtype)
+ else:
+ guide_scale = None
+
+ latent = self.diffusion.sample(
+ noise=noise,
+ sampler=solver_sample,
+ model=get_model(self.diffusion_model),
+ model_kwargs={"cond": ctx, "guidance": guide_scale},
+ steps=sample_steps,
+ show_progress=True,
+ guide_scale=guide_scale,
+ return_intermediate=None,
+ **kwargs)
+
+ self.dynamic_unload(self.diffusion_model,
+ 'diffusion_model',
+ skip_loaded=True)
+
+ if 'latent' in value_output:
+ if value_output['latent'] is None or (
+ isinstance(value_output['latent'], list)
+ and len(value_output['latent']) < 1):
+ value_output['latent'] = []
+ value_output['latent'].append(latent)
+
+ self.dynamic_load(self.first_stage_model, 'first_stage_model')
+ x_samples = self.decode_first_stage(latent).float()
+ self.dynamic_unload(self.first_stage_model,
+ 'first_stage_model',
+ skip_loaded=True)
+ images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
+ if 'images' in value_output:
+ if value_output['images'] is None or (
+ isinstance(value_output['images'], list)
+ and len(value_output['images']) < 1):
+ value_output['images'] = []
+ value_output['images'].append(images)
+
+ for k, v in value_output.items():
+ if isinstance(v, list):
+ value_output[k] = torch.cat(v, dim=0)
+ if isinstance(v, torch.Tensor):
+ value_output[k] = v.cpu()
+
+ # unregister tuner
+ if tuner_model is not None and tuner_model != '' and len(
+ tuner_model) > 0:
+ self.tuner_infer.unregister_tuner(tuner_model,
+ self.diffusion_model,
+ self.cond_stage_model)
+
+ # unregister control
+ if control_model is not None and control_model != '':
+ self.control_infer.unregister_controllers(control_model,
+ self.diffusion_model)
+
+ return value_output
diff --git a/scepter/modules/inference/largen_inference.py b/scepter/modules/inference/largen_inference.py
index 7caf409..55c43fd 100644
--- a/scepter/modules/inference/largen_inference.py
+++ b/scepter/modules/inference/largen_inference.py
@@ -40,7 +40,7 @@ class LargenInference(DiffusionInference):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_path)
else:
- sd = torch.load(local_path, map_location='cpu')
+ sd = torch.load(local_path, map_location='cpu', weights_only=True)
if 'model' in sd:
sd = sd['model']
diff --git a/scepter/modules/inference/pixart_inference.py b/scepter/modules/inference/pixart_inference.py
index fc58e24..2a72beb 100644
--- a/scepter/modules/inference/pixart_inference.py
+++ b/scepter/modules/inference/pixart_inference.py
@@ -84,9 +84,11 @@ class PixArtInference(DiffusionInference):
function_name)(value_input['prompt'],
return_mask=True)
context['crossattn'] = cont.float()
+ self.dynamic_load(self.diffusion_model, 'diffusion_model')
null_context['crossattn'] = get_model(
self.diffusion_model).y_embedder.y_embedding[None].repeat(
num_samples, 1, 1)
+ self.dynamic_unload(self.diffusion_model, 'diffusion_model')
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
diff --git a/scepter/modules/inference/tuner_inference.py b/scepter/modules/inference/tuner_inference.py
index 79fbecb..48f758d 100644
--- a/scepter/modules/inference/tuner_inference.py
+++ b/scepter/modules/inference/tuner_inference.py
@@ -13,10 +13,6 @@ try:
from peft.utils import CONFIG_NAME, SAFETENSORS_WEIGHTS_NAME, WEIGHTS_NAME
except Exception as e:
warnings.warn(f'Import peft error, please deal with this problem: {e}')
-try:
- from swift import Swift, SwiftModel
-except Exception as e:
- warnings.warn(f'Import swift error, please deal with this problem: {e}')
class TunerInference():
@@ -27,6 +23,11 @@ class TunerInference():
# @classmethod
def unregister_tuner(self, tuner_model_list, diffusion_model,
cond_stage_model):
+ try:
+ from swift import SwiftModel
+ except Exception as e:
+ warnings.warn(f'Import swift error, please deal with this problem: {e}')
+
self.logger.info('Unloading tuner model')
if isinstance(diffusion_model['model'], SwiftModel):
for adapter_name in diffusion_model['model'].adapters:
@@ -41,6 +42,11 @@ class TunerInference():
# @classmethod
def register_tuner(self, tuner_model_list, diffusion_model,
cond_stage_model):
+ try:
+ from swift import Swift
+ except Exception as e:
+ warnings.warn(f'Import swift error, please deal with this problem: {e}')
+
self.logger.info('Loading tuner model')
if len(tuner_model_list) < 1:
self.unregister_tuner(tuner_model_list, diffusion_model,
@@ -137,7 +143,7 @@ class TunerInference():
state_dict = {}
is_bin_file = True
if os.path.isfile(bin_file):
- state_dict = torch.load(bin_file)
+ state_dict = torch.load(bin_file, weights_only=True)
elif os.path.isfile(safe_file):
is_bin_file = False
from safetensors.torch import \
diff --git a/scepter/modules/model/__init__.py b/scepter/modules/model/__init__.py
index 66e7062..bf33490 100644
--- a/scepter/modules/model/__init__.py
+++ b/scepter/modules/model/__init__.py
@@ -2,4 +2,4 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model import (backbone, embedder, head, loss, metric,
- neck, network, tokenizer, tuner)
+ neck, network, tokenizer, tuner, diffusion)
diff --git a/scepter/modules/model/backbone/__init__.py b/scepter/modules/model/backbone/__init__.py
index 6f5c074..afb7fcf 100644
--- a/scepter/modules/model/backbone/__init__.py
+++ b/scepter/modules/model/backbone/__init__.py
@@ -1,4 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone import (autoencoder, image, mmdit, pixart,
- unet, utils, video)
+ unet, utils, video, flux)
diff --git a/scepter/modules/model/backbone/flux/__init__.py b/scepter/modules/model/backbone/flux/__init__.py
new file mode 100644
index 0000000..97a75f6
--- /dev/null
+++ b/scepter/modules/model/backbone/flux/__init__.py
@@ -0,0 +1 @@
+from .flux import Flux
\ No newline at end of file
diff --git a/scepter/modules/model/backbone/flux/flux.py b/scepter/modules/model/backbone/flux/flux.py
new file mode 100644
index 0000000..7265601
--- /dev/null
+++ b/scepter/modules/model/backbone/flux/flux.py
@@ -0,0 +1,245 @@
+import math
+from functools import partial
+
+import torch
+from einops import rearrange, repeat
+from scepter.modules.model.base_model import BaseModel
+from scepter.modules.model.registry import BACKBONES
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+from torch import Tensor, nn
+from torch.utils.checkpoint import checkpoint_sequential
+
+from .layers import (DoubleStreamBlock, EmbedND, LastLayer,
+ MLPEmbedder, SingleStreamBlock,
+ timestep_embedding)
+
+@BACKBONES.register_class()
+class Flux(BaseModel):
+ """
+ Transformer backbone Diffusion model with RoPE.
+ """
+ para_dict = {
+ "IN_CHANNELS": {
+ "value": 64,
+ "description": "model's input channels."
+ },
+ "OUT_CHANNELS": {
+ "value": 64,
+ "description": "model's output channels."
+ },
+ "HIDDEN_SIZE": {
+ "value": 1024,
+ "description": "model's hidden size."
+ },
+ "NUM_HEADS": {
+ "value": 16,
+ "description": "number of heads in the transformer."
+ },
+ "AXES_DIM": {
+ "value": [16, 56, 56],
+ "description": "dimensions of the axes of the positional encoding."
+ },
+ "THETA": {
+ "value": 10_000,
+ "description": "theta for positional encoding."
+ },
+ "VEC_IN_DIM": {
+ "value": 768,
+ "description": "dimension of the vector input."
+ },
+ "GUIDANCE_EMBED": {
+ "value": False,
+ "description": "whether to use guidance embedding."
+ },
+ "CONTEXT_IN_DIM": {
+ "value": 4096,
+ "description": "dimension of the context input."
+ },
+ "MLP_RATIO": {
+ "value": 4.0,
+ "description": "ratio of mlp hidden size to hidden size."
+ },
+ "QKV_BIAS": {
+ "value": True,
+ "description": "whether to use bias in qkv projection."
+ },
+ "DEPTH": {
+ "value": 19,
+ "description": "number of transformer blocks."
+ },
+ "DEPTH_SINGLE_BLOCKS": {
+ "value": 38,
+ "description": "number of transformer blocks in the single stream block."
+ },
+ "USE_GRAD_CHECKPOINT": {
+ "value": False,
+ "description": "whether to use gradient checkpointing."
+ }
+ }
+ def __init__(
+ self,
+ cfg,
+ logger = None
+ ):
+ super().__init__(cfg, logger=logger)
+ self.in_channels = cfg.IN_CHANNELS
+ self.out_channels = cfg.get("OUT_CHANNELS", self.in_channels)
+ hidden_size = cfg.get("HIDDEN_SIZE", 1024)
+ num_heads = cfg.get("NUM_HEADS", 16)
+ axes_dim = cfg.AXES_DIM
+ theta = cfg.THETA
+ vec_in_dim = cfg.VEC_IN_DIM
+ self.guidance_embed = cfg.GUIDANCE_EMBED
+ context_in_dim = cfg.CONTEXT_IN_DIM
+ mlp_ratio = cfg.MLP_RATIO
+ qkv_bias = cfg.QKV_BIAS
+ depth = cfg.DEPTH
+ depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS
+ self.use_grad_checkpoint = cfg.get("USE_GRAD_CHECKPOINT", False)
+
+ if hidden_size % num_heads != 0:
+ raise ValueError(
+ f"Hidden size {hidden_size} must be divisible by num_heads {num_heads}"
+ )
+ pe_dim = hidden_size // num_heads
+ if sum(axes_dim) != pe_dim:
+ raise ValueError(f"Got {axes_dim} but expected positional dim {pe_dim}")
+ self.hidden_size = hidden_size
+ self.num_heads = num_heads
+ self.pe_embedder = EmbedND(dim=pe_dim, theta=theta, axes_dim= axes_dim)
+ self.img_in = nn.Linear(self.in_channels, self.hidden_size, bias=True)
+ self.time_in = MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size)
+ self.vector_in = MLPEmbedder(vec_in_dim, self.hidden_size)
+ self.guidance_in = (
+ MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size) if self.guidance_embed else nn.Identity()
+ )
+ self.txt_in = nn.Linear(context_in_dim, self.hidden_size)
+
+ self.double_blocks = nn.ModuleList(
+ [
+ DoubleStreamBlock(
+ self.hidden_size,
+ self.num_heads,
+ mlp_ratio=mlp_ratio,
+ qkv_bias=qkv_bias,
+ )
+ for _ in range(depth)
+ ]
+ )
+
+ self.single_blocks = nn.ModuleList(
+ [
+ SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=mlp_ratio)
+ for _ in range(depth_single_blocks)
+ ]
+ )
+
+ self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
+
+ def prepare_input(self, x, context, y, x_shape=None):
+ # x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360]
+ bs, c, h, w = x.shape
+ x = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2)
+ x_id = torch.zeros(h // 2, w // 2, 3)
+ x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None]
+ x_id[..., 2] = x_id[..., 2] + torch.arange(w // 2)[None, :]
+ x_ids = repeat(x_id, "h w c -> b (h w) c", b=bs)
+ txt_ids = torch.zeros(bs, context.shape[1], 3)
+ return x, x_ids.to(x), context.to(x), txt_ids.to(x), y.to(x), h, w
+
+ def unpack(self, x: Tensor, height: int, width: int) -> Tensor:
+ return rearrange(
+ x,
+ "b (h w) (c ph pw) -> b c (h ph) (w pw)",
+ h=math.ceil(height/2),
+ w=math.ceil(width/2),
+ ph=2,
+ pw=2,
+ )
+
+ def load_pretrained_model(self, pretrained_model):
+ if next(self.parameters()).device.type == 'meta':
+ map_location = we.device_id
+ else:
+ map_location = "cpu"
+ if pretrained_model is not None:
+ with FS.get_from(pretrained_model, wait_finish=True) as local_model:
+ if local_model.endswith('safetensors'):
+ from safetensors.torch import load_file as load_safetensors
+ sd = load_safetensors(local_model, device=map_location)
+ else:
+ sd = torch.load(local_model, map_location=map_location)
+ missing, unexpected = self.load_state_dict(sd, strict=False, assign=True)
+ self.logger.info(
+ f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
+ )
+ if len(missing) > 0:
+ self.logger.info(f'Missing Keys:\n {missing}')
+ if len(unexpected) > 0:
+ self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
+
+ def forward(
+ self,
+ x: Tensor,
+ t: Tensor,
+ cond: dict = {},
+ guidance: Tensor | None = None,
+ gc_seg: int = 0
+ ) -> Tensor:
+ x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(x, cond["context"], cond["y"])
+ # running on sequences img
+ x = self.img_in(x)
+ vec = self.time_in(timestep_embedding(t, 256))
+ if self.guidance_embed:
+ if guidance is None:
+ raise ValueError("Didn't get guidance strength for guidance distilled model.")
+ vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
+ vec = vec + self.vector_in(y)
+ txt = self.txt_in(txt)
+ ids = torch.cat((txt_ids, x_ids), dim=1)
+ pe = self.pe_embedder(ids)
+ kwargs = dict(
+ vec=vec,
+ pe=pe,
+ txt_length=txt.shape[1],
+ )
+ x = torch.cat((txt, x), 1)
+ if self.use_grad_checkpoint and gc_seg >= 0:
+ x = checkpoint_sequential(
+ functions=[partial(block, **kwargs) for block in self.double_blocks],
+ segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
+ input=x,
+ use_reentrant=False
+ )
+ else:
+ for block in self.double_blocks:
+ x = block(x, **kwargs)
+
+ kwargs = dict(
+ vec=vec,
+ pe=pe,
+ )
+
+ if self.use_grad_checkpoint and gc_seg >= 0:
+ x = checkpoint_sequential(
+ functions=[partial(block, **kwargs) for block in self.single_blocks],
+ segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
+ input=x,
+ use_reentrant=False
+ )
+ else:
+ for block in self.single_blocks:
+ x = block(x, **kwargs)
+ x = x[:, txt.shape[1] :, ...]
+ x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
+ x = self.unpack(x, h, w)
+ return x
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('MODEL',
+ __class__.__name__,
+ Flux.para_dict,
+ set_name=True)
diff --git a/scepter/modules/model/backbone/flux/layers.py b/scepter/modules/model/backbone/flux/layers.py
new file mode 100644
index 0000000..696a40c
--- /dev/null
+++ b/scepter/modules/model/backbone/flux/layers.py
@@ -0,0 +1,282 @@
+from __future__ import annotations
+
+import math
+from dataclasses import dataclass
+from torch import Tensor, nn
+import torch
+from einops import rearrange, repeat
+from torch import Tensor
+
+
+def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
+ q, k = apply_rope(q, k, pe)
+ x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
+ x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
+ x = rearrange(x, "B H L D -> B L (H D)")
+ return x
+
+
+def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
+ assert dim % 2 == 0
+ scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
+ omega = 1.0 / (theta**scale)
+ out = torch.einsum("...n,d->...nd", pos, omega)
+ out = torch.stack([torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1)
+ out = rearrange(out, "b n d (i j) -> b n d i j", i=2, j=2)
+ return out.float()
+
+
+def apply_rope(xq: Tensor, xk: Tensor, freqs_cis: Tensor) -> tuple[Tensor, Tensor]:
+ xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
+ xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
+ xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
+ xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
+ return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
+
+class EmbedND(nn.Module):
+ def __init__(self, dim: int, theta: int, axes_dim: list[int]):
+ super().__init__()
+ self.dim = dim
+ self.theta = theta
+ self.axes_dim = axes_dim
+
+ def forward(self, ids: Tensor) -> Tensor:
+ n_axes = ids.shape[-1]
+ emb = torch.cat(
+ [rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)],
+ dim=-3,
+ )
+
+ return emb.unsqueeze(1)
+
+
+def timestep_embedding(t: Tensor, dim, max_period=10000, time_factor: float = 1000.0):
+ """
+ Create sinusoidal timestep embeddings.
+ :param t: a 1-D Tensor of N indices, one per batch element.
+ These may be fractional.
+ :param dim: the dimension of the output.
+ :param max_period: controls the minimum frequency of the embeddings.
+ :return: an (N, D) Tensor of positional embeddings.
+ """
+ t = time_factor * t
+ half = dim // 2
+ freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
+ t.device
+ )
+
+ args = t[:, None].float() * freqs[None]
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
+ if torch.is_floating_point(t):
+ embedding = embedding.to(t)
+ return embedding
+
+
+class MLPEmbedder(nn.Module):
+ def __init__(self, in_dim: int, hidden_dim: int):
+ super().__init__()
+ self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True)
+ self.silu = nn.SiLU()
+ self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True)
+
+ def forward(self, x: Tensor) -> Tensor:
+ return self.out_layer(self.silu(self.in_layer(x)))
+
+
+class RMSNorm(torch.nn.Module):
+ def __init__(self, dim: int):
+ super().__init__()
+ self.scale = nn.Parameter(torch.ones(dim))
+
+ def forward(self, x: Tensor):
+ x_dtype = x.dtype
+ x = x.float()
+ rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + 1e-6)
+ return (x * rrms).to(dtype=x_dtype) * self.scale
+
+
+class QKNorm(torch.nn.Module):
+ def __init__(self, dim: int):
+ super().__init__()
+ self.query_norm = RMSNorm(dim)
+ self.key_norm = RMSNorm(dim)
+
+ def forward(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]:
+ q = self.query_norm(q)
+ k = self.key_norm(k)
+ return q.to(v), k.to(v)
+
+
+class SelfAttention(nn.Module):
+ def __init__(self, dim: int, num_heads: int = 8, qkv_bias: bool = False):
+ super().__init__()
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.norm = QKNorm(head_dim)
+ self.proj = nn.Linear(dim, dim)
+
+ def forward(self, x: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
+ qkv = self.qkv(x)
+ q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ q, k = self.norm(q, k, v)
+ x = attention(q, k, v, pe=pe, mask=mask)
+ x = self.proj(x)
+ return x
+
+
+@dataclass
+class ModulationOut:
+ shift: Tensor
+ scale: Tensor
+ gate: Tensor
+
+
+class Modulation(nn.Module):
+ def __init__(self, dim: int, double: bool):
+ super().__init__()
+ self.is_double = double
+ self.multiplier = 6 if double else 3
+ self.lin = nn.Linear(dim, self.multiplier * dim, bias=True)
+
+ def forward(self, vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
+ out = self.lin(nn.functional.silu(vec))[:, None, :].chunk(self.multiplier, dim=-1)
+
+ return (
+ ModulationOut(*out[:3]),
+ ModulationOut(*out[3:]) if self.is_double else None,
+ )
+
+
+class DoubleStreamBlock(nn.Module):
+ def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False):
+ super().__init__()
+
+ mlp_hidden_dim = int(hidden_size * mlp_ratio)
+ self.num_heads = num_heads
+ self.hidden_size = hidden_size
+ self.img_mod = Modulation(hidden_size, double=True)
+ self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.img_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
+
+ self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.img_mlp = nn.Sequential(
+ nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
+ )
+
+ self.txt_mod = Modulation(hidden_size, double=True)
+ self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.txt_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
+
+ self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.txt_mlp = nn.Sequential(
+ nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
+ )
+
+ def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None, txt_length = None):
+ img_mod1, img_mod2 = self.img_mod(vec)
+ txt_mod1, txt_mod2 = self.txt_mod(vec)
+
+ txt, img = x[:, :txt_length], x[:, txt_length:]
+
+ # prepare image for attention
+ img_modulated = self.img_norm1(img)
+ img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
+ img_qkv = self.img_attn.qkv(img_modulated)
+ img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
+ # prepare txt for attention
+ txt_modulated = self.txt_norm1(txt)
+ txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
+ txt_qkv = self.txt_attn.qkv(txt_modulated)
+ txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
+
+ # run actual attention
+ q = torch.cat((txt_q, img_q), dim=2)
+ k = torch.cat((txt_k, img_k), dim=2)
+ v = torch.cat((txt_v, img_v), dim=2)
+ if mask is not None:
+ mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
+ attn = attention(q, k, v, pe=pe, mask = mask)
+ txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
+
+ # calculate the img bloks
+ img = img + img_mod1.gate * self.img_attn.proj(img_attn)
+ img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
+
+ # calculate the txt bloks
+ txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
+ txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
+ x = torch.cat((txt, img), 1)
+ return x
+
+
+class SingleStreamBlock(nn.Module):
+ """
+ A DiT block with parallel linear layers as described in
+ https://arxiv.org/abs/2302.05442 and adapted modulation interface.
+ """
+
+ def __init__(
+ self,
+ hidden_size: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ qk_scale: float | None = None,
+ ):
+ super().__init__()
+ self.hidden_dim = hidden_size
+ self.num_heads = num_heads
+ head_dim = hidden_size // num_heads
+ self.scale = qk_scale or head_dim**-0.5
+
+ self.mlp_hidden_dim = int(hidden_size * mlp_ratio)
+ # qkv and mlp_in
+ self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + self.mlp_hidden_dim)
+ # proj and mlp_out
+ self.linear2 = nn.Linear(hidden_size + self.mlp_hidden_dim, hidden_size)
+
+ self.norm = QKNorm(head_dim)
+
+ self.hidden_size = hidden_size
+ self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+
+ self.mlp_act = nn.GELU(approximate="tanh")
+ self.modulation = Modulation(hidden_size, double=False)
+
+ def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None) -> Tensor:
+ mod, _ = self.modulation(vec)
+ x_mod = (1 + mod.scale) * self.pre_norm(x) + mod.shift
+ qkv, mlp = torch.split(self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1)
+
+ q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ q, k = self.norm(q, k, v)
+ if mask is not None:
+ mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
+ # compute attention
+ attn = attention(q, k, v, pe=pe, mask = mask)
+ # compute activation in mlp stream, cat again and run second linear layer
+ output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
+ return x + mod.gate * output
+
+
+class LastLayer(nn.Module):
+ def __init__(self, hidden_size: int, patch_size: int, out_channels: int):
+ super().__init__()
+ self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
+ self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
+
+ def forward(self, x: Tensor, vec: Tensor) -> Tensor:
+ shift, scale = self.adaLN_modulation(vec).chunk(2, dim=1)
+ x = (1 + scale[:, None, :]) * self.norm_final(x) + shift[:, None, :]
+ x = self.linear(x)
+ return x
diff --git a/scepter/modules/model/backbone/pixart/pixart_alpha.py b/scepter/modules/model/backbone/pixart/pixart_alpha.py
index cbae899..7febcbb 100644
--- a/scepter/modules/model/backbone/pixart/pixart_alpha.py
+++ b/scepter/modules/model/backbone/pixart/pixart_alpha.py
@@ -60,8 +60,8 @@ class FinalLayer(nn.Module):
self.linear = nn.Linear(hidden_size,
patch_size * patch_size * out_channels,
bias=True)
- self.adaLN_modulation = nn.Sequential(
- nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
+ self.adaLN_modulation = nn.Sequential(nn.SiLU(),
+ nn.Linear(hidden_size, 2 * hidden_size, bias=True))
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
diff --git a/scepter/modules/model/backbone/transformer/layers.py b/scepter/modules/model/backbone/transformer/layers.py
index 10875a9..3d559a8 100644
--- a/scepter/modules/model/backbone/transformer/layers.py
+++ b/scepter/modules/model/backbone/transformer/layers.py
@@ -9,7 +9,6 @@ import math
import torch
import torch.nn as nn
from einops import rearrange
-
from scepter.modules.model.backbone.transformer.attention import drop_path
@@ -152,7 +151,6 @@ class SizeEmbedder(TimestepEmbedder):
@property
def dtype(self):
- # 返回模型参数的数据类型
return next(self.parameters()).dtype
diff --git a/scepter/modules/model/backbone/transformer/pos_embed.py b/scepter/modules/model/backbone/transformer/pos_embed.py
index d1b2f58..2d6515b 100644
--- a/scepter/modules/model/backbone/transformer/pos_embed.py
+++ b/scepter/modules/model/backbone/transformer/pos_embed.py
@@ -8,6 +8,7 @@ from itertools import repeat as iter_repeat
from typing import Iterable
import numpy as np
+
import torch
@@ -117,7 +118,6 @@ def apply_2d_rope(xq,
# xq_.shape = [b, seq_len, dim // 2, 2]
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 2)
- # 转为复数域
xq_ = torch.view_as_complex(
xq_) # [b, seq_len, dim // 2, 2]=>xq.shape = [b, seq_len, dim]
xk_ = torch.view_as_complex(xk_)
diff --git a/scepter/modules/model/base_model.py b/scepter/modules/model/base_model.py
index 6685a2a..f98a5e0 100644
--- a/scepter/modules/model/base_model.py
+++ b/scepter/modules/model/base_model.py
@@ -43,7 +43,13 @@ class BaseModel(nn.Module):
self._dist_data[key][k] += v
else:
self._dist_data[key][k] = v
-
+ def collect_probe(self):
+ probe_data_dict = self._probe_data
+ for k, v in self._modules.items():
+ if isinstance(getattr(self, k), BaseModel):
+ for kk, vv in getattr(self, k).collect_probe().items():
+ probe_data_dict[f'{k}/{kk}'] = vv
+ return probe_data_dict
def probe_data(self):
gather_probe_data = gather_data(self._probe_data)
_dist_data_list = gather_data([self._dist_data])
diff --git a/scepter/modules/model/diffusion/__init__.py b/scepter/modules/model/diffusion/__init__.py
new file mode 100644
index 0000000..42baa1c
--- /dev/null
+++ b/scepter/modules/model/diffusion/__init__.py
@@ -0,0 +1,3 @@
+from .samplers import BaseDiffusionSampler, FlowEluerSampler, DDIMSampler
+from .schedules import BaseNoiseScheduler, ScaledLinearScheduler, FlowMatchShiftScheduler
+from .diffusions import BaseDiffusion, DiffusionFluxRF
\ No newline at end of file
diff --git a/scepter/modules/model/diffusion/diffusions.py b/scepter/modules/model/diffusion/diffusions.py
new file mode 100644
index 0000000..3fdf31a
--- /dev/null
+++ b/scepter/modules/model/diffusion/diffusions.py
@@ -0,0 +1,264 @@
+import os
+import math
+import torch
+from collections import OrderedDict
+
+from scepter.modules.utils.config import dict_to_yaml, Config
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+from scepter.modules.model.registry import DIFFUSIONS, NOISE_SCHEDULERS, DIFFUSION_SAMPLERS
+from tqdm import trange
+
+@DIFFUSIONS.register_class()
+class BaseDiffusion(object):
+ para_dict = {
+ "NOISE_SCHEDULER": {},
+ "SAMPLER_SCHEDULER": {},
+ "MIN_SNR_GAMMA": {
+ "value": None,
+ "description": "The minimum SNR gamma value for the loss function."
+ },
+ "PREDICTION_TYPE": {
+ "value": "eps",
+ "description": "The type of prediction to use for the loss function."
+ }
+ }
+ def __init__(self, cfg, logger=None):
+ super(BaseDiffusion, self).__init__()
+ self.logger = logger
+ self.cfg = cfg
+ self.init_params()
+
+ def init_params(self):
+ self.min_snr_gamma = self.cfg.get("MIN_SNR_GAMMA", None)
+ self.prediction_type = self.cfg.get("PREDICTION_TYPE", "eps")
+ self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER, logger=self.logger)
+ self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get("SAMPLER_SCHEDULER", self.cfg.NOISE_SCHEDULER),
+ logger=self.logger)
+ self.num_timesteps = self.noise_scheduler.num_timesteps
+ if self.cfg.have("WORK_DIR") and we.rank == 0:
+ schedule_visualization = os.path.join(self.cfg.WORK_DIR, "noise_schedule.png")
+ with FS.put_to(schedule_visualization) as local_path:
+ self.noise_scheduler.plot_noise_sampling_map(local_path)
+ schedule_visualization = os.path.join(self.cfg.WORK_DIR, "sampler_schedule.png")
+ with FS.put_to(schedule_visualization) as local_path:
+ self.sampler_scheduler.plot_noise_sampling_map(local_path)
+
+
+ def sample(self, noise, model, model_kwargs={}, steps=20, sampler=None, use_dynamic_cfg=False, guide_scale=None, guide_rescale=None,
+ show_progress=False, return_intermediate=None, intermediate_callback=None):
+ assert isinstance(steps, (int, torch.LongTensor))
+ assert return_intermediate in (None, 'x0', 'xt')
+ assert isinstance(sampler, (str, dict, Config))
+ intermediates = []
+
+ def callback_fn(x_t, t, sigma=None, alpha=None):
+ t = t.repeat(len(x_t)).round().long().to(x_t.device)
+ if guide_scale is None or guide_scale == 1.0:
+ out = model(x=x_t, t=t, **model_kwargs)
+ else:
+ if use_dynamic_cfg:
+ guidance_scale = 1 + guide_scale * ((1 - math.cos(math.pi * ((steps - t.item()) / steps) ** 5.0)) / 2)
+ else:
+ guidance_scale = guide_scale
+ y_out = model(x=x_t, t=t, **model_kwargs[0])
+ u_out = model(x=x_t, t=t, **model_kwargs[1])
+ out = u_out + guidance_scale * (y_out - u_out)
+ if guide_rescale is not None and guide_rescale > 0.0:
+ ratio = (
+ y_out.flatten(1).std(dim=1) /
+ (out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
+ (y_out.ndim - 1))
+ out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
+
+ if self.prediction_type == 'x0':
+ x0 = out
+ elif self.prediction_type == 'eps':
+ x0 = (x_t - sigma * out) / alpha
+ elif self.prediction_type == 'v':
+ x0 = alpha * x_t - sigma * out
+ else:
+ raise NotImplementedError(
+ f'prediction_type {self.prediction_type} not implemented')
+
+ # print("torch.sum(y_out):", torch.sum(y_out), "torch.sum(u_out):", torch.sum(u_out), "torch.sum(out):",
+ # torch.sum(out), "torch.sum(x0):", torch.sum(x0), "sigmas", sigma, "alphas", alpha)
+ return x0
+
+ sampler_ins = self.get_sampler(sampler)
+
+ # this is ignored for schnell
+ sampler_output = sampler_ins.preprare_sampler(
+ noise,
+ steps=steps,
+ prediction_type=self.prediction_type,
+ scheduler_ins=self.sampler_scheduler,
+ callback_fn=callback_fn
+ )
+
+ for _ in trange(steps, disable=not show_progress):
+ trange.desc = sampler_output.msg
+ sampler_output = sampler_ins.step(sampler_output)
+ if return_intermediate == 'x_0':
+ intermediates.append(sampler_output.x_0)
+ elif return_intermediate == 'x_t':
+ intermediates.append(sampler_output.x_t)
+ if intermediate_callback is not None:
+ intermediate_callback(intermediates[-1])
+ return (sampler_output.x_0, intermediates) if return_intermediate is not None else sampler_output.x_0
+
+
+ def loss(self, x_0, model, model_kwargs={}, reduction='mean', noise=None, **kwargs):
+ # use noise scheduler to add noise
+ if noise is None:
+ noise = torch.randn_like(x_0)
+ schedule_output = self.noise_scheduler.add_noise(x_0, noise)
+ x_t, t, sigma, alpha = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha
+ out = model(x=x_t, t=t, **model_kwargs)
+
+ # mse loss
+ target = {
+ 'eps': noise,
+ 'x0': x_0,
+ 'v': alpha * noise - sigma * x_0
+ }[self.prediction_type]
+
+ loss = (out - target).pow(2)
+ if reduction == 'mean':
+ loss = loss.flatten(1).mean(dim=1)
+
+ if self.min_snr_gamma is not None:
+ alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
+ sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
+ snrs = (alphas / sigmas).clamp(min=1e-20)
+ min_snrs = snrs.clamp(max=self.min_snr_gamma)
+ weights = min_snrs / snrs
+ else:
+ weights = 1
+
+ loss = loss * weights
+ return loss
+
+ def get_sampler(self, sampler):
+ if isinstance(sampler, str):
+ if sampler not in DIFFUSION_SAMPLERS.class_map:
+ if self.logger is not None:
+ self.logger.info(f"{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}")
+ else:
+ print(f"{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}")
+ return None
+ sampler_cfg = Config(cfg_dict={"NAME": sampler}, load=False)
+ sampler_ins = DIFFUSION_SAMPLERS.build(sampler_cfg, logger=self.logger)
+ elif isinstance(sampler, (Config, dict, OrderedDict)):
+ if isinstance(sampler, (dict, OrderedDict)):
+ sampler = Config(cfg_dict={k.upper():v for k, v in dict(sampler).items()}, load=False)
+ sampler_ins = DIFFUSION_SAMPLERS.build(sampler, logger=self.logger)
+ else:
+ raise NotImplementedError
+ return sampler_ins
+
+ def __repr__(self) -> str:
+ return f'{self.__class__.__name__}' + ' ' + super().__repr__()
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('DIFFUSIONS',
+ __class__.__name__,
+ BaseDiffusion.para_dict,
+ set_name=True)
+
+@DIFFUSIONS.register_class()
+class DiffusionFluxRF(BaseDiffusion):
+ para_dict = {
+ "PREDICTION_TYPE": {
+ "value": "raw",
+ "description": "The type of prediction to use for the loss function."
+ }
+ }
+ para_dict.update(BaseDiffusion.para_dict)
+ def __init__(self, cfg, logger=None):
+ super(DiffusionFluxRF, self).__init__(cfg, logger=logger)
+ self.prediction_type = self.cfg.get("PREDICTION_TYPE", "raw")
+
+ def loss(self, x_0, model, model_kwargs={}, reduction='mean', noise=None, **kwargs):
+ if noise is None:
+ noise = torch.randn_like(x_0)
+ schedule_output = self.noise_scheduler.add_noise(x_0, noise)
+ x_t, t, sigma = schedule_output.x_t, schedule_output.t, schedule_output.sigma
+ out = model(x=x_t, t=sigma, **model_kwargs)
+ # raw
+ if self.prediction_type == "raw":
+ target = noise - x_0
+ out = out
+ elif self.prediction_type == "sigma_scaled":
+ target = x_0
+ out = out * (-sigma) + x_t
+ else:
+ raise NotImplementedError
+
+ loss = (target - out) ** 2
+ if reduction == 'mean':
+ loss = loss.flatten(1).mean(dim=1)
+
+ if self.min_snr_gamma is not None:
+ alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
+ sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
+ snrs = (alphas / sigmas).clamp(min=1e-20)
+ min_snrs = snrs.clamp(max=self.min_snr_gamma)
+ weights = min_snrs / snrs
+ else:
+ weights = 1
+
+ loss = loss * weights
+ return loss
+
+ @torch.no_grad()
+ def sample(self,
+ noise,
+ model,
+ model_kwargs={},
+ steps=20,
+ sampler = None,
+ show_progress=False,
+ return_intermediate=None,
+ intermediate_callback=None,
+ **kwargs):
+ # sanity check
+ assert isinstance(steps, (int, torch.LongTensor))
+ assert return_intermediate in (None, 'x0', 'xt')
+ assert isinstance(sampler, (str, dict, Config))
+ intermediates = []
+
+ def callback_fn(x_t, t, sigma):
+ sigma = torch.full((x_t.shape[0],), sigma, dtype=x_t.dtype, device=x_t.device)
+ x_0 = model(x = x_t, t=sigma, **model_kwargs)
+ return x_0
+
+ sampler_ins = self.get_sampler(sampler)
+
+ # this is ignored for schnell
+ sampler_output = sampler_ins.preprare_sampler(
+ noise,
+ steps = steps,
+ prediction_type=self.prediction_type,
+ scheduler_ins = self.sampler_scheduler,
+ callback_fn=callback_fn
+ )
+
+ for _ in trange(steps, disable=not show_progress):
+ trange.desc = sampler_output.msg
+ sampler_output = sampler_ins.step(sampler_output)
+ if return_intermediate == 'x_0':
+ intermediates.append(sampler_output.x_0)
+ elif return_intermediate == 'x_t':
+ intermediates.append(sampler_output.x_t)
+ if intermediate_callback is not None:
+ intermediate_callback(intermediates[-1])
+ return (sampler_output.x_0, intermediates) if return_intermediate is not None else sampler_output.x_t
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('DIFFUSIONS',
+ __class__.__name__,
+ DiffusionFluxRF.para_dict,
+ set_name=True)
\ No newline at end of file
diff --git a/scepter/modules/model/diffusion/samplers.py b/scepter/modules/model/diffusion/samplers.py
new file mode 100644
index 0000000..0237b32
--- /dev/null
+++ b/scepter/modules/model/diffusion/samplers.py
@@ -0,0 +1,211 @@
+from dataclasses import dataclass, field
+import torch
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.model.registry import DIFFUSION_SAMPLERS
+from .util import _i
+
+@dataclass
+class SamplerOutput(object):
+ callback_fn: callable
+ prediction_type: str
+ alphas: torch.Tensor
+ betas: torch.Tensor
+ sigmas: torch.Tensor
+ alphas_init: torch.Tensor
+ betas_init: torch.Tensor
+ sigmas_init: torch.Tensor
+ ts: torch.Tensor
+ x_t: torch.Tensor
+ x_0: torch.Tensor
+ step: int
+ msg: str
+
+ def add_custom_field(self, key: str, value) -> None:
+ self.__setattr__(key, value)
+
+
+@DIFFUSION_SAMPLERS.register_class("base")
+class BaseDiffusionSampler(object):
+ para_dict = {}
+ def __init__(self, cfg, logger=None):
+ super(BaseDiffusionSampler, self).__init__()
+ self.logger = logger
+ self.cfg = cfg
+ self.init_params()
+
+ def init_params(self):
+ self.discretization_type = self.cfg.get("DISCRETIZATION_TYPE", "linspace")
+ self.discard_penultimate_step = self.cfg.get("DISCARD_PENULTIMATE_STEP", False)
+ self.free_steps = self.cfg.get("FREE_STEPS", None)
+ self.t_max = self.cfg.get("T_MAX", None)
+ self.t_min = self.cfg.get("T_MIN", None)
+
+ def discretization(self, steps=20, num_timesteps=1000, **kwargs):
+ # get timesteps
+ if isinstance(steps, int):
+ steps += 1 if self.discard_penultimate_step else 0
+ t_max = num_timesteps - 1 if self.t_max is None else self.t_max
+ t_min = 0 if self.t_min is None else self.t_min
+
+ # discretize timesteps
+ if self.discretization_type == 'leading':
+ steps = torch.arange(t_min, t_max + 1,
+ (t_max - t_min + 1) / steps).flip(0)
+ elif self.discretization_type == 'linspace':
+ steps = torch.linspace(t_max, t_min, steps)
+ elif self.discretization_type == 'trailing':
+ steps = torch.arange(t_max, t_min - 1,
+ -((t_max - t_min + 1) / steps))
+ elif self.discretization_type == 'free':
+ steps = torch.tensor(self.free_steps)
+ else:
+ raise NotImplementedError(
+ f'{self.discretization_type} discretization not implemented')
+ steps = steps.clamp_(t_min, t_max)
+ elif isinstance(steps, list):
+ steps = torch.tensor(steps)
+ timesteps = torch.as_tensor(steps, dtype=torch.float32)
+ return timesteps
+
+ def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
+ sigmas=None, betas=None, alphas=None, callback_fn = None,
+ **kwargs):
+ '''
+ 1. Control the model's inputs and outputs externally in the solver by callback_fn,
+ and perform the conversion between x0 and xt internally within the solver.
+ 2. The function callback_fn use the x0 and xt as the default inputs and also give me
+ the x0 and xt as output. The other inputs will be set in kwargs.
+ 3. The basic parameters of the schedule should be set manually.
+ 4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
+ which manage all necessary information.
+ '''
+ num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
+ timestamps = self.discretization(steps, num_timesteps=num_timesteps, **kwargs)
+ alphas = scheduler_ins.t_to_alpha(timestamps, **kwargs) if scheduler_ins is not None else alphas
+ betas = scheduler_ins.t_to_beta(timestamps, **kwargs) if scheduler_ins is not None else betas
+ sigmas = scheduler_ins.t_to_sigma(timestamps, **kwargs) if scheduler_ins is not None else sigmas
+ alphas_init = scheduler_ins.t_to_alpha_init(timestamps, **kwargs) if scheduler_ins is not None else alphas
+ betas_init = scheduler_ins.t_to_beta_init(timestamps, **kwargs) if scheduler_ins is not None else betas
+ sigmas_init = scheduler_ins.t_to_sigma_init(timestamps, **kwargs) if scheduler_ins is not None else sigmas
+
+ output = SamplerOutput(
+ callback_fn=callback_fn,
+ prediction_type=prediction_type,
+ alphas=alphas,
+ betas=betas,
+ sigmas=sigmas,
+ alphas_init=alphas_init,
+ betas_init=betas_init,
+ sigmas_init=sigmas_init,
+ ts=timestamps,
+ x_t=noise,
+ x_0=noise,
+ step=0,
+ msg=f"step 0"
+ )
+ return output
+
+ def step(self, sampler_ouput):
+ raise NotImplementedError(f'DiffusionSampler step function not implemented')
+
+ def __repr__(self) -> str:
+ return f'{self.__class__.__name__}' + ' ' + super().__repr__()
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('DIFFUSION_SAMPLERS',
+ __class__.__name__,
+ BaseDiffusionSampler.para_dict,
+ set_name=True)
+
+
+@DIFFUSION_SAMPLERS.register_class("eluer")
+class EulerSampler(BaseDiffusionSampler):
+ def step(self, sampler_ouput):
+ pass
+
+
+@DIFFUSION_SAMPLERS.register_class("ddim")
+class DDIMSampler(BaseDiffusionSampler):
+
+ def init_params(self):
+ super().init_params()
+ self.eta = self.cfg.get('ETA', 0.)
+ self.discretization_type = self.cfg.get("DISCRETIZATION_TYPE", "trailing")
+
+ def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
+ sigmas=None, betas=None, alphas=None, callback_fn = None,
+ **kwargs):
+ output = super().preprare_sampler(noise, steps, scheduler_ins, prediction_type, sigmas, betas, alphas, callback_fn, **kwargs)
+ sigmas = output.sigmas
+ sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
+ sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
+ sigmas_vp[sigmas == float('inf')] = 1.
+ output.add_custom_field('sigmas_vp', sigmas_vp)
+ return output
+
+ def step(self, sampler_output):
+ step = sampler_output.step
+ x_t = sampler_output.x_t
+ t = sampler_output.ts[step]
+ sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
+ alpha_init = _i(sampler_output.alphas_init, step, x_t)
+ sigma_init = _i(sampler_output.sigmas_init, step, x_t)
+
+ x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
+ noise_factor = self.eta * (sigmas_vp[step + 1] ** 2 / sigmas_vp[step] ** 2 *
+ (1 - (1 - sigmas_vp[step] ** 2) /
+ (1 - sigmas_vp[step + 1] ** 2)))
+ d = (x_t - (1 - sigmas_vp[step] ** 2) ** 0.5 * x) / sigmas_vp[step]
+ x = (1 - sigmas_vp[step + 1] ** 2) ** 0.5 * x + \
+ (sigmas_vp[step + 1] ** 2 - noise_factor ** 2) ** 0.5 * d
+ sampler_output.x_0 = x
+ if sigmas_vp[step + 1] > 0:
+ x += noise_factor * torch.randn_like(x)
+ # print("i:", step, "sigma_init:", sigma_init, "alpha_init", alpha_init, "sigmas_vp[i]", sigmas_vp[step], "torch.sum(x_0):", torch.sum(x_0), "torch.sum(x):", torch.sum(x))
+ sampler_output.x_t = x
+ sampler_output.step += 1
+ sampler_output.msg = f'step {step}'
+ return sampler_output
+
+
+@DIFFUSION_SAMPLERS.register_class("flow_eluer")
+class FlowEluerSampler(BaseDiffusionSampler):
+ def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
+ sigmas=None, betas=None, alphas=None, callback_fn = None,
+ **kwargs):
+ if noise.ndim == 3:
+ seq_len = noise.shape[2] // 4
+ else:
+ n, _, h, w = noise.shape
+ seq_len = (h // 2 * w // 2)
+ kwargs["seq_len"] = seq_len
+ output = super().preprare_sampler(noise, steps, scheduler_ins, prediction_type, sigmas, betas, alphas, callback_fn, **kwargs)
+ return output
+
+ def step(self, sampler_output):
+ step = sampler_output.step
+ x_t = sampler_output.x_t
+ sigma_curr, sigma_prev = sampler_output.sigmas[step], sampler_output.sigmas[step + 1]
+ prediction_type = sampler_output.prediction_type
+ assert prediction_type in ("raw", "sigma_scaled")
+ t = sampler_output.ts[step]
+ x_0 = sampler_output.callback_fn(x_t, t, sigma_curr)
+ x_t = x_t + (sigma_prev - sigma_curr) * x_0
+ sampler_output.x_0 = x_0
+ sampler_output.x_t = x_t
+ sampler_output.step += 1
+ sampler_output.msg = f'step {step}, sigma_curr: {sigma_curr}, sigma_prev: {sigma_prev}'
+ return sampler_output
+
+ def discretization(self, steps=20, num_timesteps = 1000, **kwargs):
+ # extra step for zero
+ timesteps = torch.linspace(num_timesteps, 0, steps + 1)
+ return timesteps
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('DIFFUSION_SAMPLERS',
+ __class__.__name__,
+ FlowEluerSampler.para_dict,
+ set_name=True)
\ No newline at end of file
diff --git a/scepter/modules/model/diffusion/schedules.py b/scepter/modules/model/diffusion/schedules.py
new file mode 100644
index 0000000..898532d
--- /dev/null
+++ b/scepter/modules/model/diffusion/schedules.py
@@ -0,0 +1,533 @@
+import math
+from dataclasses import dataclass, field
+from typing import Callable
+
+import torch
+import numpy as np
+
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.math_plot import plot_multi_curves
+from scepter.modules.model.registry import NOISE_SCHEDULERS
+from torch import Tensor
+
+from .util import _i
+
+@dataclass
+class ScheduleOutput(object):
+ x_t: torch.Tensor
+ x_0: torch.Tensor
+ t: torch.Tensor
+ sigma: torch.Tensor
+ alpha: torch.Tensor
+ custom_fields: dict = field(default_factory=dict)
+
+ def add_custom_field(self, key: str, value) -> None:
+ self.__setattr__(key, value)
+
+
+@NOISE_SCHEDULERS.register_class()
+class BaseNoiseScheduler(object):
+ '''
+ In the diffusion model, the parameters related to the noise schedule are alpha, beta,
+ and sigma. The following are the definitions of the above three parameters, which should
+ be the basic property for the instance of noise scheduler.
+ \alpha_{t} = \sqrt{1 - \beta_{t}^2} \alpha is the strength of signal and \beta is the strength of noise
+ \sigma_{t} = \sqrt{1 - \overline\alpha} = \sqrt{1 - \prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
+
+ where sigma_{t} is the var of p(x_{t-1}|x_{t}, x_{0}).
+
+ (reference to https://arxiv.org/abs/2010.02502)
+ let sigma transfer to beta:
+ square_\beta = 1 - \frac{1 - square_\sigma_{t}}{1 - square_\sigma_{t - 1 }}
+
+ '''
+ para_dict = {
+ "NUM_TIMESTEPS": {
+ "value": 1000,
+ "description": "The number of timesteps for sampling."
+ },
+ }
+ def __init__(self, cfg, logger=None):
+ super(BaseNoiseScheduler, self).__init__()
+ self.logger = logger
+ self.cfg = cfg
+ self.init_params()
+ self.get_schedule()
+ # self.check_function()
+
+ def init_params(self):
+ self.num_timesteps = self.cfg.get("NUM_TIMESTEPS", 1000)
+ self._sample_steps = torch.arange(self.num_timesteps, dtype=torch.float32)
+ self._sigmas, self._betas, self._alphas, self._timesteps = None, None, None, None
+ def check_function(self):
+ # for the same t, we should gurantee t_to_sigma and sigma_to_t is aligned
+ try:
+ predict_timestamps = self.sigma_to_t(self.sigmas)
+ predict_sigmas = self.t_to_sigma(self._timesteps)
+ diff_sigmas = torch.sum(torch.abs(predict_sigmas - self.sigmas))
+ diff_timestamps = torch.sum(torch.abs(predict_timestamps - self._timesteps))
+ if diff_sigmas > 1e-3 or diff_timestamps > 1:
+ self.logger.info(f"The noise scheduler {self.__class__.__name__} is not correct, "
+ f"please check the function sigma_to_t or t_to_sigma."
+ f"Info: diff sigmas {diff_sigmas}, diff timestamps {diff_timestamps}")
+ raise "The noise scheduler checked failed."
+ else:
+ self.logger.info(f"The noise scheduler {self.__class__.__name__} is checked and passed.")
+ except Exception as e:
+ if isinstance(e, NotImplementedError):
+ self.logger.info("Not implemented function sigma_to_t or t_to_sigma, skip check.")
+ else:
+ self.logger.info(f"The noise scheduler {self.__class__.__name__} is not correct, "
+ f"please check the function sigma_to_t or t_to_sigma. Error: {e}")
+ raise e
+
+ def get_schedule(self):
+ raise NotImplementedError(f'NoiseScheduler get_schedule function not implemented')
+
+ def square_betas_to_sigmas(self, square_betas):
+ return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
+ def sigmas_to_square_betas(self, sigmas):
+ square_alphas = 1 - sigmas ** 2
+ betas = 1 - torch.cat([square_alphas[:1], square_alphas[1:] / square_alphas[:-1]])
+ return betas
+ def sigma_to_t(self, sigma, **kwargs):
+ if sigma == float('inf'):
+ t = torch.full_like(sigma, len(self._sigmas) - 1)
+ else:
+ log_sigmas = torch.sqrt(self._sigmas**2 /
+ (1 - self._sigmas**2)).log().to(sigma)
+ log_sigma = sigma.log()
+ dists = log_sigma - log_sigmas[:, None]
+ low_idx = dists.ge(0).cumsum(dim=0).argmax(dim=0).clamp(
+ max=log_sigmas.shape[0] - 2)
+ high_idx = low_idx + 1
+ low, high = log_sigmas[low_idx], log_sigmas[high_idx]
+ w = (low - log_sigma) / (low - high)
+ w = w.clamp(0, 1)
+ t = (1 - w) * low_idx + w * high_idx
+ t = t.view(sigma.shape)
+ if t.ndim == 0:
+ t = t.unsqueeze(0)
+ return t
+ def t_to_sigma(self, t, **kwargs):
+ t = t.float()
+ low_idx, high_idx, w = t.floor().long(), t.ceil().long(), t.frac()
+ log_sigmas = torch.sqrt(self.sigmas**2 /
+ (1 - self.sigmas**2)).log().to(t)
+ log_sigma = (1 - w) * log_sigmas[low_idx] + w * log_sigmas[high_idx]
+ log_sigma[torch.isnan(log_sigma)
+ | torch.isinf(log_sigma)] = float('inf')
+ return log_sigma.exp()
+
+ def t_to_alpha(self, t, **kwargs):
+ sigma = self.t_to_sigma(t)
+ square_beta = self.sigmas_to_square_betas(sigma)
+ return torch.sqrt(1 - square_beta)
+
+ def t_to_beta(self, t, **kwargs):
+ sigma = self.t_to_sigma(t)
+ square_beta = self.sigmas_to_square_betas(sigma)
+ return torch.sqrt(square_beta)
+
+ def add_noise(self, x_0, noise = None, t = None):
+ if t is None:
+ t = torch.randint(0, self.num_timesteps, (x_0.shape[0],), device=x_0.device).long()
+ alpha = _i(self.alphas, t, x_0)
+ sigma = _i(self.sigmas, t, x_0)
+ x_t = alpha * x_0 + sigma * noise
+
+ return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, alpha=alpha, sigma=sigma)
+
+ def t_to_alpha_init(self, t, **kwargs):
+ indices = t.long()
+ indices[indices >= self.num_timesteps] = self.num_timesteps - 1
+ timesteps = self.timesteps.to(t)[indices]
+ step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
+ alpha = self.alphas[step_indices].flatten().to(t)
+ return alpha
+
+ def t_to_beta_init(self, t, **kwargs):
+ indices = t.long()
+ indices[indices >= self.num_timesteps] = self.num_timesteps - 1
+ timesteps = self.timesteps.to(t)[indices]
+ step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
+ beta = self.betas[step_indices].flatten().to(t)
+ return beta
+
+ def t_to_sigma_init(self, t, **kwargs):
+ indices = t.long()
+ indices[indices >= self.num_timesteps] = self.num_timesteps - 1
+ timesteps = self.timesteps.to(t)[indices]
+ step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
+ sigma = self.sigmas[step_indices].flatten().to(t)
+ return sigma
+
+ def rescale_zero_terminal_snr(self, alphas_cumprod):
+ """
+ Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
+ Args:
+ betas (`torch.Tensor`):
+ the betas that the scheduler is being initialized with.
+ Returns:
+ `torch.Tensor`: rescaled betas with zero terminal SNR
+ """
+ alphas_bar_sqrt = alphas_cumprod.sqrt()
+ # Store old values.
+ alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
+ alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
+ # Shift so the last timestep is zero.
+ alphas_bar_sqrt -= alphas_bar_sqrt_T
+ # Scale so the first timestep is back to the old value.
+ alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
+ # Convert alphas_bar_sqrt to betas
+ alphas_bar = alphas_bar_sqrt ** 2 # Revert sqrt
+ return alphas_bar
+
+ @property
+ def sigmas(self):
+ return self._sigmas
+
+ @property
+ def betas(self):
+ return self._betas
+
+ @property
+ def alphas(self):
+ return self._alphas
+
+ @property
+ def timesteps(self):
+ return self._timesteps
+
+ # plot the noise sampling map
+ def plot_noise_sampling_map(self, save_path):
+ y = [
+ {"data": self._sigmas.cpu().numpy(), "label": "sigmas"},
+ {"data": self._betas.cpu().numpy(), "label": "betas"},
+ {"data": self._alphas.cpu().numpy(), "label": "alphas"},
+ {"data": self._timesteps.cpu().numpy()/self.num_timesteps, "label": "timesteps"}
+ ]
+ plot_multi_curves(
+ x=self._sample_steps.cpu().numpy(),
+ y=y,
+ x_label='timesteps',
+ y_label=None,
+ title=f"{self.__class__.__name__}'s noise sampling map",
+ save_path=save_path
+ )
+ return save_path
+
+ def __repr__(self) -> str:
+ return f'{self.__class__.__name__}' + ' ' + super().__repr__()
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('NOISE_SCHEDULER',
+ __class__.__name__,
+ BaseNoiseScheduler.para_dict,
+ set_name=True)
+
+
+@NOISE_SCHEDULERS.register_class()
+class ScaledLinearScheduler(BaseNoiseScheduler):
+ para_dict = {}
+ def init_params(self):
+ super().init_params()
+ self.beta_min = self.cfg.get('BETA_MIN', 0.00085)
+ self.beta_max = self.cfg.get('BETA_MAX', 0.012)
+ self.snr_shift_scale = self.cfg.get('SNR_SHIFT_SCALE', None)
+ self.rescale_betas_zero_snr = self.cfg.get('RESCALE_BETAS_ZERO_SNR', False)
+
+ def square_betas_to_sigmas(self, square_betas, snr_shift_scale=None, rescale_betas_zero_snr=False):
+ if snr_shift_scale is not None or rescale_betas_zero_snr:
+ alphas_cumprod = torch.cumprod(1 - square_betas, dim=0)
+ if snr_shift_scale is not None and snr_shift_scale > 0:
+ alphas_cumprod = alphas_cumprod / (snr_shift_scale + (1 - snr_shift_scale) * alphas_cumprod)
+ if rescale_betas_zero_snr:
+ alphas_cumprod = self.rescale_zero_terminal_snr(alphas_cumprod)
+ return torch.sqrt(1 - alphas_cumprod)
+ else:
+ return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
+
+ def get_schedule(self):
+ square_betas = torch.linspace(self.beta_min**0.5, self.beta_max**0.5, self.num_timesteps, dtype=torch.float32) ** 2
+ self._sigmas = self.square_betas_to_sigmas(square_betas, self.snr_shift_scale, self.rescale_betas_zero_snr)
+ self._betas = torch.sqrt(square_betas)
+ self._alphas = torch.sqrt(1 - self._sigmas ** 2)
+ self._timesteps = torch.arange(len(self._sigmas), dtype=torch.float32)
+
+
+@NOISE_SCHEDULERS.register_class()
+class FlowMatchUniformScheduler(BaseNoiseScheduler):
+ def get_schedule(self):
+ timesteps = np.linspace(1, self.num_timesteps, self.num_timesteps, dtype=np.float32).copy()
+ timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
+ self._timesteps = timesteps
+ self._sigmas = self.t_to_sigma(timesteps)
+ self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
+ self._alphas = torch.sqrt(1 - self.betas ** 2)
+
+ def add_noise(self, x_0, noise = None, t = None):
+ if t is None:
+ t = torch.rand((x_0.shape[0],), device=x_0.device)
+ sigma = self.t_to_sigma(t)
+ shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
+ x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
+ return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, sigma=sigma, alpha=self.t_to_alpha(t))
+
+ def sigma_to_t(self, sigma, **kwargs):
+ return sigma * self.num_timesteps
+
+ def t_to_sigma(self, t, **kwargs):
+ return t/self.num_timesteps
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('NOISE_SCHEDULER',
+ __class__.__name__,
+ FlowMatchUniformScheduler.para_dict,
+ set_name=True)
+
+
+@NOISE_SCHEDULERS.register_class()
+class FlowMatchSigmoidScheduler(FlowMatchUniformScheduler):
+ para_dict = {
+ "SIGMOID_SCALE": {
+ "value": 1,
+ "description": "The scale for the sigmoid function."
+ }
+ }
+ def init_params(self):
+ super().init_params()
+ self.sigmoid_scale = self.cfg.get("SIGMOID_SCALE", 1)
+
+ def sigma_to_t(self, sigma, **kwargs):
+ t = - torch.log(1/sigma - 1)/self.sigmoid_scale
+ return t * self.num_timesteps
+
+ def t_to_sigma(self, t, **kwargs):
+ return torch.sigmoid(self.sigmoid_scale * t/self.num_timesteps)
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('NOISE_SCHEDULER',
+ __class__.__name__,
+ FlowMatchSigmoidScheduler.para_dict,
+ set_name=True)
+
+@NOISE_SCHEDULERS.register_class()
+class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
+ para_dict = {
+ "SHIFT": {
+ "value": 3,
+ "description": "The shift factor for the timestamp."
+ },
+ "SIGMOID_SCALE": {
+ "value": 1,
+ "description": "The scale for the sigmoid function."
+ }
+ }
+
+ def init_params(self):
+ super().init_params()
+ self.shift = self.cfg.get("SHIFT", 3)
+ self.sigmoid_scale = self.cfg.get("SIGMOID_SCALE", 1)
+
+ def add_noise(self, x_0, noise = None, t = None):
+ if t is None:
+ logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
+ logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
+ t = logits_norm.sigmoid() * self.num_timesteps
+ sigma = self.t_to_sigma(t)
+ shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
+ x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
+ return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, sigma=sigma, alpha=self.t_to_alpha(t))
+
+ def sigma_to_t(self, sigma, **kwargs):
+ t = sigma/(sigma - self.shift * sigma + self.shift)
+ return t * self.num_timesteps
+
+ def t_to_sigma(self, t, **kwargs):
+ t = t / self.num_timesteps
+ return (t * self.shift) / (1 + (self.shift - 1) * t)
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('NOISE_SCHEDULER',
+ __class__.__name__,
+ FlowMatchShiftScheduler.para_dict,
+ set_name=True)
+
+@NOISE_SCHEDULERS.register_class()
+class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
+ para_dict = {
+ "SHIFT": {
+ "value": True,
+ "description": "Use timestamp shift or not, default is True."
+ },
+ "SIGMOID_SCALE": {
+ "value": 1,
+ "description": "The scale of sigmoid function for sampling timesteps."
+ },
+ "BASE_SHIFT": {
+ "value": 0.5,
+ "description": "The base shift factor for the timestamp."
+ },
+ "MAX_SHIFT": {
+ "value": 1.15,
+ "description": "The max shift factor for the timestamp."
+ }
+ }
+
+ def init_params(self):
+ super().init_params()
+ self.shift = self.cfg.get("SHIFT", True)
+ self.sigmoid_scale = self.cfg.get("SIGMOID_SCALE", 1)
+ self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
+ self.max_shift = self.cfg.get('MAX_SHIFT', 1.15)
+
+ def time_shift(self, mu: float, sigma_scale: float, t: Tensor):
+ return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma_scale)
+ def sigma_shift(self, mu: float, sigma_scale: float, sigma: Tensor):
+ return 1/(torch.pow((1-sigma) * math.exp(mu)/sigma, sigma_scale) + 1)
+
+ def get_lin_function(self,
+ x1: float = 256, y1: float = 0.5, x2: float = 4096, y2: float = 1.15
+ ) -> Callable[[float], float]:
+ m = (y2 - y1) / (x2 - x1)
+ b = y1 - m * x1
+ return lambda x: m * x + b
+
+ def add_noise(self, x_0, noise = None, t = None):
+ if x_0.ndim == 3:
+ seq_len = x_0.shape[2] // 4
+ else:
+ n, _, h, w = x_0.shape
+ seq_len = (h // 2 * w // 2)
+ if t is None:
+ logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
+ logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
+ t = logits_norm.sigmoid() * self.num_timesteps
+ sigma = self.t_to_sigma(t, seq_len=seq_len)
+ shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
+ x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
+ return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, sigma=sigma, alpha=self.t_to_alpha(t))
+
+ def sigma_to_t(self, sigma, **kwargs):
+ seq_len = kwargs.get('seq_len', 256)
+ if self.shift:
+ mu = self.get_lin_function(y1=self.base_shift, y2=self.max_shift)(seq_len)
+ sigma = self.sigma_shift(mu, 1.0, sigma)
+ t = torch.as_tensor(sigma, dtype=torch.float32)
+ return t * self.num_timesteps
+
+ def t_to_sigma(self, t, **kwargs):
+ seq_len = kwargs.get('seq_len', 256)
+ t = t/self.num_timesteps
+ if self.shift:
+ mu = self.get_lin_function(y1=self.base_shift, y2=self.max_shift)(seq_len)
+ t = self.time_shift(mu, 1.0, t)
+ sigma = torch.as_tensor(t, dtype=torch.float32)
+ return sigma
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('NOISE_SCHEDULER',
+ __class__.__name__,
+ FlowMatchFluxShiftScheduler.para_dict,
+ set_name=True)
+
+@NOISE_SCHEDULERS.register_class()
+class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
+ para_dict = {
+ "WEIGHTING_SCHEME" : {
+ "value": "logit_normal",
+ "description": "The weighting scheme for sampling timesteps, "
+ "choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']."
+ },
+ "SHIFT": {
+ "value": 3.0,
+ "description": "The shift factor for the timestamp."
+ },
+ "LOGIT_MEAN" : {
+ "value": 0.0,
+ "description": "The mean of the logit distribution for sampling timesteps."
+ },
+ "LOGIT_STD" : {
+ "value": 1.0,
+ "description": "The standard deviation of the logit distribution for sampling timesteps."
+ },
+ "MODE_SCALE" : {
+ "value": 1.29,
+ "description": "The scale factor for the mode of the logit distribution for sampling timesteps."
+ }
+ }
+
+ def init_params(self):
+ super().init_params()
+ self.weighting_scheme = self.cfg.get("WEIGHTING_SCHEME", "logit_normal")
+ self.logit_mean = self.cfg.get("LOGIT_MEAN", 0.0)
+ self.logit_std = self.cfg.get("LOGIT_STD", 1.0)
+ self.mode_scale = self.cfg.get("MODE_SCALE", 1.29)
+ self.shift = self.cfg.get("SHIFT", 1.0)
+
+ def get_schedule(self):
+ timesteps = np.linspace(1, self.num_timesteps, self.num_timesteps, dtype=np.float32).copy()
+ timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
+ self._timesteps = timesteps
+ timesteps = timesteps / self.num_timesteps
+ self._sigmas = self.shift * timesteps / (1 + (self.shift - 1) * timesteps)
+ self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
+ self._alphas = torch.sqrt(1 - self.betas ** 2)
+
+ def add_noise(self, x_0, noise=None, t=None):
+ if t is None:
+ if self.weighting_scheme == "logit_normal":
+ t = torch.normal(mean=self.logit_mean, std=self.logit_std, size=(x_0.shape[0],), device=x_0.device)
+ else:
+ t = torch.rand(x_0.shape[0], device=x_0.device)
+ t = self.compute_density_for_timestep_sampling(t) * self.num_timesteps
+ sigma = self.t_to_sigma(t)
+ shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
+ x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
+ return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, sigma=sigma, alpha=self.t_to_alpha(t))
+
+ def compute_density_for_timestep_sampling(self, t):
+ """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 self.weighting_scheme == "logit_normal":
+ # See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
+ t = torch.nn.functional.sigmoid(t)
+ elif self.weighting_scheme == "mode":
+ t = 1 - t - self.mode_scale * (torch.cos(math.pi * t / 2) ** 2 - 1 + t)
+ return t
+
+ def sigma_to_t(self, sigma, **kwargs):
+ raise NotImplementedError
+
+ def t_to_sigma(self, t, **kwargs):
+ indices = t.long()
+ indices[indices >= self.num_timesteps] = self.num_timesteps - 1
+ timesteps = self.timesteps.to(t)[indices]
+ step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
+ sigma = self.sigmas[step_indices].flatten().to(t)
+ return sigma
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('NOISE_SCHEDULER',
+ __class__.__name__,
+ FlowMatchSigmaScheduler.para_dict,
+ set_name=True)
+
+
+if __name__ == '__main__':
+ from scepter.modules.utils.config import Config
+ cfg = Config(cfg_dict={
+ "NAME": "FlowMatchShiftScheduler",
+ "SHIFT": 1.15
+ }, load=False)
+
+ scheduler = NOISE_SCHEDULERS.build(cfg)
diff --git a/scepter/modules/model/diffusion/util.py b/scepter/modules/model/diffusion/util.py
new file mode 100644
index 0000000..d6a1cad
--- /dev/null
+++ b/scepter/modules/model/diffusion/util.py
@@ -0,0 +1,10 @@
+import torch
+
+def _i(tensor, t, x):
+ """
+ Index tensor using t and format the output according to x.
+ """
+ shape = (x.size(0),) + (1,) * (x.ndim - 1)
+ if isinstance(t, torch.Tensor):
+ t = t.to(tensor.device)
+ return tensor[t].view(shape).to(x.device)
\ No newline at end of file
diff --git a/scepter/modules/model/embedder/__init__.py b/scepter/modules/model/embedder/__init__.py
index cca51d2..88524d8 100644
--- a/scepter/modules/model/embedder/__init__.py
+++ b/scepter/modules/model/embedder/__init__.py
@@ -5,3 +5,4 @@ from scepter.modules.model.embedder.embedder import (
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2,
FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner,
IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF)
+from scepter.modules.model.embedder.flux_embedder import HFEmbedder
diff --git a/scepter/modules/model/embedder/flux_embedder.py b/scepter/modules/model/embedder/flux_embedder.py
new file mode 100644
index 0000000..4e0cc15
--- /dev/null
+++ b/scepter/modules/model/embedder/flux_embedder.py
@@ -0,0 +1,163 @@
+import torch
+from scepter.modules.model.embedder.base_embedder import BaseEmbedder
+from scepter.modules.model.registry import EMBEDDERS
+from scepter.modules.model.tokenizer.tokenizer_component import whitespace_clean, basic_clean, canonicalize
+from scepter.modules.utils.config import dict_to_yaml
+import transformers
+from scepter.modules.utils.file_system import FS
+
+
+@EMBEDDERS.register_class()
+class HFEmbedder(BaseEmbedder):
+ para_dict = {
+ "HF_MODEL_CLS": {
+ "value": None,
+ "description": "huggingface cls in transfomer"
+ },
+ "MODEL_PATH": {
+ "value": None,
+ "description": "model folder path"
+ },
+ "HF_TOKENIZER_CLS": {
+ "value": None,
+ "description": "huggingface cls in transfomer"
+ },
+
+ "TOKENIZER_PATH": {
+ "value": None,
+ "description": "tokenizer folder path"
+ },
+ "MAX_LENGTH": {
+ "value": 77,
+ "description": "max length of input"
+ },
+ "OUTPUT_KEY": {
+ "value": "last_hidden_state",
+ "description": "output key"
+ },
+ "D_TYPE": {
+ "value": "float",
+ "description": "dtype"
+ },
+ "BATCH_INFER": {
+ "value": False,
+ "description": "batch infer"
+ }
+ }
+ para_dict.update(BaseEmbedder.para_dict)
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ hf_model_cls = cfg.get('HF_MODEL_CLS', None)
+ model_path = cfg.get("MODEL_PATH", None)
+ hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None)
+ tokenizer_path = cfg.get('TOKENIZER_PATH', None)
+ self.max_length = cfg.get('MAX_LENGTH', 77)
+ self.output_key = cfg.get("OUTPUT_KEY", "last_hidden_state")
+ self.d_type = cfg.get("D_TYPE", "float")
+ self.clean = cfg.get("CLEAN", "whitespace")
+ self.batch_infer = cfg.get("BATCH_INFER", False)
+ torch_dtype = getattr(torch, self.d_type)
+
+ assert hf_model_cls is not None and hf_tokenizer_cls is not None
+ assert model_path is not None and tokenizer_path is not None
+
+ with FS.get_dir_to_local_dir(tokenizer_path, wait_finish=True) as local_path:
+ self.tokenizer = getattr(transformers, hf_tokenizer_cls).from_pretrained(local_path,
+ max_length = self.max_length,
+ torch_dtype = torch_dtype)
+
+ with FS.get_dir_to_local_dir(model_path, wait_finish=True) as local_path:
+ self.hf_module = getattr(transformers, hf_model_cls).from_pretrained(local_path, torch_dtype = torch_dtype)
+
+
+ self.hf_module = self.hf_module.eval().requires_grad_(False)
+
+ def forward(self, text: list[str], return_mask = False):
+ batch_encoding = self.tokenizer(
+ text,
+ truncation=True,
+ max_length=self.max_length,
+ return_length=False,
+ return_overflowing_tokens=False,
+ padding="max_length",
+ return_tensors="pt",
+ )
+
+ outputs = self.hf_module(
+ input_ids=batch_encoding["input_ids"].to(self.hf_module.device),
+ attention_mask=None,
+ output_hidden_states=False,
+ )
+ if return_mask:
+ return outputs[self.output_key], batch_encoding['attention_mask'].to(self.hf_module.device)
+ else:
+ return outputs[self.output_key], None
+
+ def encode(self, text, return_mask = False):
+ if isinstance(text, str):
+ text = [text]
+ if self.clean:
+ text = [self._clean(u) for u in text]
+ if not self.batch_infer:
+ cont, mask = [], []
+ for tt in text:
+ one_cont, one_mask = self([tt], return_mask=return_mask)
+ cont.append(one_cont)
+ mask.append(one_mask)
+ if return_mask:
+ return torch.cat(cont, dim=0), torch.cat(mask, dim=0)
+ else:
+ return torch.cat(cont, dim=0)
+ else:
+ ret_data = self(text, return_mask = return_mask)
+ if return_mask:
+ return ret_data
+ else:
+ return ret_data[0]
+
+
+ def _clean(self, text):
+ if self.clean == 'whitespace':
+ text = whitespace_clean(basic_clean(text))
+ elif self.clean == 'lower':
+ text = whitespace_clean(basic_clean(text)).lower()
+ elif self.clean == 'canonicalize':
+ text = canonicalize(basic_clean(text))
+ return text
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('EMBEDDER',
+ __class__.__name__,
+ HFEmbedder.para_dict,
+ set_name=True)
+
+@EMBEDDERS.register_class()
+class T5PlusClipFluxEmbedder(BaseEmbedder):
+ """
+ Uses the OpenCLIP transformer encoder for text
+ """
+ para_dict = {
+ 'T5_MODEL': {},
+ 'CLIP_MODEL': {}
+ }
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger)
+ self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger)
+
+ def encode(self, text):
+ t5_embeds = self.t5_model.encode(text, return_mask = False)
+ clip_embeds = self.clip_model.encode(text, return_mask = False)
+ # change embedding strategy here
+ return {
+ 'context': t5_embeds,
+ 'y': clip_embeds,
+ }
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('EMBEDDER',
+ __class__.__name__,
+ T5PlusClipFluxEmbedder.para_dict,
+ set_name=True)
\ No newline at end of file
diff --git a/scepter/modules/model/network/__init__.py b/scepter/modules/model/network/__init__.py
index c9e9b23..3ab7bf7 100644
--- a/scepter/modules/model/network/__init__.py
+++ b/scepter/modules/model/network/__init__.py
@@ -5,4 +5,5 @@ from scepter.modules.model.network.classifier import Classifier
from scepter.modules.model.network.diffusion import (diffusion, schedules,
solvers)
from scepter.modules.model.network.ldm import (ldm, ldm_edit, ldm_pixart,
- ldm_sce, ldm_sd3, ldm_xl)
+ ldm_sce, ldm_sd3, ldm_xl,
+ ldm_flux)
\ No newline at end of file
diff --git a/scepter/modules/model/network/autoencoder/ae_kl.py b/scepter/modules/model/network/autoencoder/ae_kl.py
index 71bcbae..3d0fc95 100644
--- a/scepter/modules/model/network/autoencoder/ae_kl.py
+++ b/scepter/modules/model/network/autoencoder/ae_kl.py
@@ -1,18 +1,20 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+import copy
+import numbers
import re
from collections import OrderedDict
import numpy as np
import torch
-
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import BACKBONES, LOSSES, MODELS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
-
-
+from einops import repeat
+import math
+import torch.nn.functional as F
class DiagonalGaussianDistribution(object):
def __init__(self, mean, logvar, deterministic=False):
self.mean = mean
@@ -77,6 +79,11 @@ class AutoencoderKL(TrainModule):
'value': 16,
'description': ''
},
+ 'SCALE_FACTOR': {
+ 'value': None,
+ 'description':
+ 'if is not None, will used to scale the latent space.'
+ },
}
def __init__(self, cfg, logger=None):
@@ -89,6 +96,7 @@ class AutoencoderKL(TrainModule):
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
self.batch_size = self.cfg.get('BATCH_SIZE', 16)
self.use_conv = self.cfg.get('USE_CONV', True)
+ self.scale_factor = self.cfg.get('SCALE_FACTOR', None)
self.construct_network()
self.init_network()
@@ -169,6 +177,7 @@ class AutoencoderKL(TrainModule):
return z
def _encode(self, x, return_mom=False):
+
h = self.encoder(x)
moments = self.conv1(h)
if return_mom:
@@ -176,9 +185,15 @@ class AutoencoderKL(TrainModule):
mean, logvar = torch.chunk(moments, 2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
z = posterior.sample()
+ if self.scale_factor is not None and isinstance(
+ self.scale_factor, numbers.Number):
+ z = self.scale_factor * z
return z
def _decode(self, z):
+ if self.scale_factor is not None and isinstance(
+ self.scale_factor, numbers.Number):
+ z = z / self.scale_factor
z = self.conv2(z)
dec = self.decoder(z)
return dec
@@ -268,13 +283,274 @@ class AutoencoderKL(TrainModule):
AutoencoderKL.para_dict,
set_name=True)
+def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
+ """
+ Create sinusoidal timestep embeddings.
+ :param timesteps: a 1-D Tensor of N indices, one per batch element.
+ These may be fractional.
+ :param dim: the dimension of the output.
+ :param max_period: controls the minimum frequency of the embeddings.
+ :return: an [N x dim] Tensor of positional embeddings.
+ """
+ if not repeat_only:
+ half = dim // 2
+ freqs = torch.exp(
+ -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
+ ).to(device=timesteps.device)
+ args = torch.mm(timesteps.float().unsqueeze(1), freqs.unsqueeze(0)).view(timesteps.shape[0], len(freqs))
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
+ else:
+ embedding = repeat(timesteps, 'b -> b d', d=dim)
+ return embedding
-if __name__ == '__main__':
- import argparse
- from scepter.modules.utils.config import Config
- from scepter.modules.utils.logger import get_logger
- std_logger = get_logger(name='scepter')
- parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
- cfg = Config(load=True, parser_ins=parser)
- model = AutoencoderKL(cfg, logger=std_logger)
- model.load_pretrained_model(cfg.PRETRAINED_MODEL)
+@MODELS.register_class()
+class AutoencoderKLFlux(TrainModule):
+ para_dict = {
+ "ENCODER": {},
+ "DECODER": {},
+ "LOSS": {},
+ "EMBED_DIM": {
+ "value": 4,
+ "description": ""
+ },
+ "PRETRAINED_MODEL": {
+ "value": None,
+ "description": ""
+ },
+ "IGNORE_KEYS": {
+ "value": [],
+ "description": ""
+ },
+ "BATCH_SIZE": {
+ "value": 16,
+ "description": ""
+ },
+ }
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.encoder_cfg = self.cfg.ENCODER
+ self.decoder_cfg = self.cfg.DECODER
+ self.loss_cfg = self.cfg.get("LOSS", None)
+ self.embed_dim = self.cfg.get("EMBED_DIM", 4)
+ self.pretrained_model = self.cfg.get("PRETRAINED_MODEL", None)
+ #
+ self.ignore_keys = self.cfg.get("IGNORE_KEYS", [])
+ self.batch_size = self.cfg.get("BATCH_SIZE", 16)
+ self.resize_nx = self.cfg.get("RESIZE_NX", 1)
+ self.use_rembed = self.cfg.get("USE_REMBED", True)
+ self.use_conv = self.cfg.get('USE_CONV', True)
+ self.scale_factor = self.cfg.get('SCALE_FACTOR', None)
+ self.shift_factor = self.cfg.get('SHIFT_FACTOR', None)
+ self.construct_network()
+
+ def construct_network(self):
+ z_channels = self.encoder_cfg.Z_CHANNELS
+ self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger)
+ self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger)
+ self.conv1 = torch.nn.Conv2d(2 * z_channels, 2 * self.embed_dim, 1) if self.use_conv else torch.nn.Identity()
+ self.conv2 = torch.nn.Conv2d(self.embed_dim, z_channels,1) if self.use_conv else torch.nn.Identity()
+
+ # freeze encoder
+ for param in self.encoder.parameters():
+ param.requires_grad = False
+ for param in self.conv1.parameters():
+ param.requires_grad = False
+ for param in self.conv2.parameters():
+ param.requires_grad = False
+
+ def load_pretrained_model(self, pretrained_model):
+ if pretrained_model is not None:
+ with FS.get_from(pretrained_model, wait_finish=True) as local_model:
+ self.init_from_ckpt(local_model)
+ def init_from_ckpt(self, path, ignore_keys=list()):
+ if path.find('.safetensors') > -1:
+ from safetensors import safe_open
+ sd = OrderedDict()
+ with safe_open(path, framework="pt", device='cpu') as f:
+ for k in f.keys():
+ sd[k] = f.get_tensor(k)
+ else:
+ sd = torch.load(path, map_location="cpu")
+ if path.find('.pt') > -1 and 'state_dict' in sd:
+ sd = sd['state_dict']
+ elif path.find('.ckpt') > -1 and 'state_dict' in sd:
+ sd = sd['state_dict']
+ elif path.find('.pth') > -1 and 'model' in sd:
+ sd = sd['model']
+
+ new_sd = OrderedDict()
+ for k, v in sd.items():
+ ignored = False
+ for ik in ignore_keys:
+ if ik in k:
+ if we.rank == 0:
+ self.logger.info("ignore key {} from state_dict.".format(k))
+ ignored = True
+ break
+ k = k.replace("post_quant_conv", "conv2") if "post_quant_conv" in k else k
+ k = k.replace("quant_conv", "conv1") if "quant_conv" in k else k
+ if not ignored:
+ new_sd[k] = v
+
+ missing, unexpected = self.load_state_dict(new_sd, strict=False)
+ if we.rank == 0:
+ self.logger.info(f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys")
+ if len(missing) > 0:
+ self.logger.info(f"Missing Keys:\n {missing}")
+ if len(unexpected) > 0:
+ self.logger.info(f"\nUnexpected Keys:\n {unexpected}")
+
+ @torch.no_grad()
+ def encode(self, x, sample_posterior = True):
+ h = self.encoder(x)
+ moments = self.conv1(h)
+ mean, logvar = torch.chunk(moments, 2, dim=1)
+ posterior = DiagonalGaussianDistribution(mean, logvar)
+ if sample_posterior:
+ z = posterior.sample()
+ else:
+ z = posterior.mode()
+ if self.shift_factor is not None and isinstance(self.shift_factor, numbers.Number):
+ z = z - self.shift_factor
+ if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
+ z = self.scale_factor * z
+ return z, posterior
+
+ def decode(self, z, **kwargs):
+ b, c, h, w = z.size()
+ if kwargs.get('resize_nx', None) is not None and self.use_rembed:
+ resize_nx = kwargs['resize_nx']
+ if not torch.is_tensor(resize_nx):
+ resize_nx = torch.full((b,), resize_nx, device=we.device_id, dtype=z.dtype)
+ rembed = timestep_embedding(resize_nx, dim=self.decoder_cfg.CH_MULT[-1] * self.decoder_cfg.CH)
+ else:
+ rembed = None
+ if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
+ z = z / self.scale_factor
+ if self.shift_factor is not None and isinstance(self.shift_factor, numbers.Number):
+ z = z + self.shift_factor
+ z = self.conv2(z)
+ # add grad;
+ if rembed is not None:
+ dec = self.decoder(z, rembed)
+ else:
+ dec = self.decoder(z)
+ return dec
+
+ def share_forward(self, image=None, sample_posterior=True, **kwargs):
+ # rembed: resize embedding
+ if image is not None:
+ z, posterior = self.encode(image, sample_posterior = sample_posterior)
+ else:
+ latent = kwargs.pop("latent", None)
+ assert latent is not None
+ z = latent
+ posterior = None
+ if self.shift_factor is not None and isinstance(self.scale_factor, numbers.Number):
+ z = z - self.shift_factor
+ if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
+ z = self.scale_factor * z
+ dec = self.decode(z, **kwargs)
+ return dec, posterior
+
+ def forward(self, **kwargs):
+ if self.training:
+ ret = self.forward_train(**kwargs)
+ else:
+ ret = self.forward_test(**kwargs)
+ return ret
+
+ def forward_train(self,
+ image=None,
+ gt_image=None,
+ sample_posterior=True,
+ optimizer_idx=0,
+ global_step=0,
+ **kwargs):
+
+ if gt_image is None:
+ gt_image = copy.deepcopy(image)
+ reconstructions, posterior = self.share_forward(image, sample_posterior, **kwargs)
+
+ ret = {}
+ if optimizer_idx == 0:
+ # train encoder+decoder+logvar
+ aeloss, log_dict_ae = self.loss(gt_image, reconstructions, posterior, optimizer_idx,
+ global_step, last_layer=self.get_last_layer(), split="train")
+ # self.logger.info(f"aeloss: {aeloss.detach().cpu().item()}, ")
+ ret["loss"] = aeloss
+ ret.update(log_dict_ae)
+
+ if optimizer_idx == 1:
+ # train the discriminator
+ discloss, log_dict_disc = self.loss(gt_image, reconstructions, posterior, optimizer_idx,
+ global_step, last_layer=self.get_last_layer(), split="train")
+ # self.logger.info(f"discloss: {discloss.detach().cpu().item()}, ")
+ ret["loss"] = discloss
+ ret.update(log_dict_disc)
+
+ return ret
+
+ @torch.no_grad()
+ def forward_test(self,
+ image=None,
+ gt_image=None,
+ sample_posterior=True,
+ **kwargs):
+ resize_nx_ = 1
+ if image is not None:
+ b, c, h, w = image.size()
+ if kwargs.get('resize_ex'):
+ resize_nx_ = kwargs.pop('resize_ex')
+ image = F.interpolate(image, (int(float(h) / resize_nx_), int(float(w) / resize_nx_)), mode='bicubic')
+ image = F.interpolate(image, (h, w), mode='bicubic')
+ kwargs["resize_nx"] = resize_nx_
+ elif kwargs.get('resize_nx', None) is not None:
+ resize_nx_ = kwargs['resize_nx']
+
+ if gt_image is None:
+ if image is not None:
+ gt_image = copy.deepcopy(image)
+
+ # kwargs["resize_nx"] = resize_nx_
+ reconstructions, posterior = self.share_forward(image, sample_posterior, **kwargs)
+ reconstructions = torch.clamp((reconstructions + 1.0) / 2.0, min=0.0, max=1.0)
+ if gt_image is not None:
+ gt_image = torch.clamp((gt_image + 1.0) / 2.0, min=0.0, max=1.0)
+ else:
+ gt_image = [None for _ in range(reconstructions.shape[0])]
+ if image is not None:
+ lr_image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)
+ else:
+ lr_image = [None for _ in range(reconstructions.shape[0])]
+
+ ret = list()
+ if torch.is_tensor(resize_nx_):
+ resize_nx = [nx.item() for nx, in zip(resize_nx_.cpu())]
+ else:
+ resize_nx = [resize_nx_ for _ in range(reconstructions.size(0))]
+
+ for img, ori, gt_img, nx in zip(reconstructions, lr_image, gt_image, resize_nx):
+ ret.append({
+ "prompt": "",
+ "n_prompt": "",
+ "image": img,
+ "lr_image": ori,
+ "gt_image": gt_img,
+ "resize_nx": nx
+ })
+
+ return ret
+
+ def get_last_layer(self):
+ if hasattr(self.decoder, 'conv_out'):
+ return self.decoder.conv_out.weight
+ else:
+ return self.decoder.head[-1].weight
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml("MODEL", __class__.__name__, AutoencoderKLFlux.para_dict, set_name=True)
diff --git a/scepter/modules/model/network/ldm/ldm.py b/scepter/modules/model/network/ldm/ldm.py
index 2a519dd..5c98b2b 100644
--- a/scepter/modules/model/network/ldm/ldm.py
+++ b/scepter/modules/model/network/ldm/ldm.py
@@ -10,7 +10,7 @@ import torch
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.network.train_module import TrainModule
-from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES,
+from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES, DIFFUSIONS,
MODELS, TOKENIZERS)
from scepter.modules.model.utils.basic_utils import count_params, default
from scepter.modules.utils.config import dict_to_yaml
@@ -112,13 +112,18 @@ class LatentDiffusion(TrainModule):
if self.zero_terminal_snr:
assert self.parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
- self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
- n=self.num_timesteps,
- zero_terminal_snr=self.zero_terminal_snr,
- **self.schedule_args)
-
- self.diffusion = GaussianDiffusion(
- sigmas=self.sigmas, prediction_type=self.parameterization)
+ diffusion_cfg = self.cfg.get("DIFFUSION", None)
+ if diffusion_cfg is not None:
+ if self.cfg.have("WORK_DIR"):
+ diffusion_cfg.WORK_DIR = self.cfg.WORK_DIR
+ self.diffusion = DIFFUSIONS.build(diffusion_cfg, logger=self.logger)
+ else:
+ self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
+ n=self.num_timesteps,
+ zero_terminal_snr=self.zero_terminal_snr,
+ **self.schedule_args)
+ self.diffusion = GaussianDiffusion(
+ sigmas=self.sigmas, prediction_type=self.parameterization)
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
@@ -131,6 +136,7 @@ class LatentDiffusion(TrainModule):
self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
self.size_factor = self.cfg.get('SIZE_FACTOR', 8)
+ self.decoder_bias = self.cfg.get("DECODER_BIAS", 0)
self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
self.p_zero = self.cfg.get('P_ZERO', 0.0)
@@ -140,6 +146,7 @@ class LatentDiffusion(TrainModule):
if self.train_n_prompt is None:
self.train_n_prompt = ''
self.use_ema = self.cfg.get('USE_EMA', False)
+ self.eval_ema = self.cfg.get('EVAL_EMA', False)
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
def construct_network(self):
@@ -290,6 +297,7 @@ class LatentDiffusion(TrainModule):
@torch.no_grad()
@torch.autocast('cuda', dtype=torch.float16)
def forward_test(self,
+ image=None,
prompt=None,
n_prompt=None,
sampler='ddim',
diff --git a/scepter/modules/model/network/ldm/ldm_edit.py b/scepter/modules/model/network/ldm/ldm_edit.py
index 742f4b5..e43db34 100644
--- a/scepter/modules/model/network/ldm/ldm_edit.py
+++ b/scepter/modules/model/network/ldm/ldm_edit.py
@@ -111,6 +111,7 @@ class LatentDiffusionEdit(LatentDiffusion):
@torch.no_grad()
@torch.autocast('cuda', dtype=torch.float16)
def forward_test(self,
+ image=None,
prompt=None,
n_prompt=None,
sampler='ddim',
diff --git a/scepter/modules/model/network/ldm/ldm_flux.py b/scepter/modules/model/network/ldm/ldm_flux.py
new file mode 100644
index 0000000..a004edb
--- /dev/null
+++ b/scepter/modules/model/network/ldm/ldm_flux.py
@@ -0,0 +1,220 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import copy
+import math
+import numbers
+import random
+import torch
+from scepter.modules.model.network.ldm import LatentDiffusion
+from scepter.modules.model.registry import MODELS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS, DIFFUSIONS
+from scepter.modules.model.utils.basic_utils import disabled_train
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.model.utils.basic_utils import count_params
+
+
+@MODELS.register_class()
+class LatentDiffusionFlux(LatentDiffusion):
+ para_dict = LatentDiffusion.para_dict
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.guide_scale = cfg.get('GUIDE_SCALE', 3.5)
+
+ def init_params(self):
+ self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
+ assert self.parameterization in [
+ 'eps', 'x0', 'v', 'rf'
+ ], 'currently only supporting "eps" and "x0" and "v" and "rf"'
+
+ diffusion_cfg = self.cfg.get("DIFFUSION", None)
+ assert diffusion_cfg is not None
+ if self.cfg.have("WORK_DIR"):
+ diffusion_cfg.WORK_DIR = self.cfg.WORK_DIR
+ self.diffusion = DIFFUSIONS.build(diffusion_cfg, logger=self.logger)
+
+ self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
+ self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
+
+ self.model_config = self.cfg.DIFFUSION_MODEL
+ self.first_stage_config = self.cfg.FIRST_STAGE_MODEL
+ self.cond_stage_config = self.cfg.COND_STAGE_MODEL
+ self.tokenizer_config = self.cfg.get('TOKENIZER', None)
+ self.loss_config = self.cfg.get('LOSS', None)
+
+ self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
+ self.size_factor = self.cfg.get('SIZE_FACTOR', 16)
+ self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
+ self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
+ self.p_zero = self.cfg.get('P_ZERO', 0.0)
+ self.train_n_prompt = self.cfg.get('TRAIN_N_PROMPT', '')
+ if self.default_n_prompt is None:
+ self.default_n_prompt = ''
+ if self.train_n_prompt is None:
+ self.train_n_prompt = ''
+ self.use_ema = self.cfg.get('USE_EMA', False)
+ self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
+
+ def construct_network(self):
+ # embedding_context = torch.device("meta") if self.model_config.get("PRETRAINED_MODEL", None) else nullcontext()
+ # with embedding_context:
+ self.model = BACKBONES.build(self.model_config, logger=self.logger).to(torch.bfloat16)
+ self.logger.info('all parameters:{}'.format(count_params(self.model)))
+ if self.use_ema:
+ if self.model_ema_config:
+ self.model_ema = BACKBONES.build(self.model_ema_config,
+ logger=self.logger)
+ else:
+ self.model_ema = copy.deepcopy(self.model)
+ self.model_ema = self.model_ema.eval()
+ for param in self.model_ema.parameters():
+ param.requires_grad = False
+ if self.loss_config:
+ self.loss = LOSSES.build(self.loss_config, logger=self.logger)
+ if self.tokenizer_config is not None:
+ self.tokenizer = TOKENIZERS.build(self.tokenizer_config,
+ logger=self.logger)
+ if self.first_stage_config:
+ self.first_stage_model = MODELS.build(self.first_stage_config,
+ logger=self.logger)
+ self.first_stage_model = self.first_stage_model.eval()
+ self.first_stage_model.train = disabled_train
+ for param in self.first_stage_model.parameters():
+ param.requires_grad = False
+ else:
+ self.first_stage_model = None
+ if self.tokenizer_config is not None:
+ self.cond_stage_config.KWARGS = {
+ 'vocab_size': self.tokenizer.vocab_size
+ }
+ if self.cond_stage_config == '__is_unconditional__':
+ print(
+ f'Training {self.__class__.__name__} as an unconditional model.'
+ )
+ self.cond_stage_model = None
+ else:
+ model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger)
+ self.cond_stage_model = model.eval().requires_grad_(False)
+ self.cond_stage_model.train = disabled_train
+
+ def noise_sample(self, num_samples, h, w, seed, dtype = torch.bfloat16):
+ noise = torch.randn(
+ num_samples,
+ 16,
+ # allow for packing
+ 2 * math.ceil(h / 16),
+ 2 * math.ceil(w / 16),
+ device=we.device_id,
+ dtype=dtype,
+ generator=torch.Generator(device=we.device_id).manual_seed(seed),
+ )
+ return noise
+
+ def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
+ x_start = self.encode_first_stage(image, **kwargs)
+ if prompt and self.cond_stage_model:
+ ctx = getattr(self.cond_stage_model, 'encode')(prompt)
+ else:
+ assert False
+ if 'index' in kwargs:
+ kwargs.pop('index')
+ guide_scale = self.guide_scale
+ if guide_scale is not None:
+ guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype)
+ else:
+ guide_scale = None
+ loss = self.diffusion.loss(x_0=x_start,
+ model=self.model,
+ model_kwargs={"cond": ctx, "guidance": guide_scale},
+ noise=noise,
+ **kwargs)
+ loss = loss.mean()
+ ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
+ return ret
+
+
+ @torch.no_grad()
+ def forward_test(self,
+ image=None,
+ prompt=None,
+ sampler='flow_eluer',
+ sample_steps=20,
+ seed=2023,
+ guide_scale=4.5,
+ guide_rescale=0.0,
+ show_process=False,
+ **kwargs):
+ seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
+ if isinstance(prompt, str):
+ prompt = [prompt]
+ assert isinstance(prompt, list)
+ num_samples = len(prompt)
+ if prompt and self.cond_stage_model:
+ ctx = getattr(self.cond_stage_model, 'encode')(prompt)
+ else:
+ assert False
+ if 'index' in kwargs:
+ kwargs.pop('index')
+ image_size = None
+ if 'meta' in kwargs:
+ meta = kwargs.pop('meta')
+ if 'image_size' in meta:
+ h = int(meta['image_size'][0][0])
+ w = int(meta['image_size'][1][0])
+ image_size = [h, w]
+ if 'image_size' in kwargs:
+ image_size = kwargs.pop('image_size')
+ if isinstance(image_size, numbers.Number):
+ image_size = [image_size, image_size]
+ if image_size is None:
+ image_size = [1024, 1024]
+ height, width = image_size
+ noise = self.noise_sample(
+ num_samples,
+ height,
+ width,
+ seed
+ )
+ guide_scale = guide_scale or self.guide_scale
+ if guide_scale is not None:
+ guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device, dtype=noise.dtype)
+ else:
+ guide_scale = None
+ # UNet use input n_prompt
+ samples = self.diffusion.sample(
+ noise=noise,
+ sampler=sampler,
+ model=self.model,
+ model_kwargs= {"cond": ctx, "guidance": guide_scale},
+ steps=sample_steps,
+ show_progress=True,
+ guide_scale = guide_scale,
+ return_intermediate=None,
+ **kwargs).float()
+ with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
+ x_samples = self.decode_first_stage(samples).float()
+ x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
+ outputs = list()
+
+ for i, (p, img) in enumerate(zip(prompt, x_samples)):
+ one_tup = {'prompt': str(p), 'n_prompt': '', 'image': img}
+ outputs.append(one_tup)
+
+ return outputs
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('MODEL',
+ __class__.__name__,
+ LatentDiffusionFlux.para_dict,
+ set_name=True)
+
+ @torch.no_grad()
+ def encode_first_stage(self, x, **kwargs):
+ z = self.first_stage_model.encode(x)
+ if isinstance(z, (tuple, list)):
+ z = z[0]
+ return z
+
+ @torch.no_grad()
+ def decode_first_stage(self, z):
+ return self.first_stage_model.decode(z)
diff --git a/scepter/modules/model/network/ldm/ldm_pixart.py b/scepter/modules/model/network/ldm/ldm_pixart.py
index e338b73..fd20b9a 100644
--- a/scepter/modules/model/network/ldm/ldm_pixart.py
+++ b/scepter/modules/model/network/ldm/ldm_pixart.py
@@ -134,6 +134,7 @@ class LatentDiffusionPixart(LatentDiffusion):
@torch.no_grad()
def forward_test(self,
+ image=None,
prompt=None,
label=None,
sampler='ddim',
diff --git a/scepter/modules/model/network/ldm/ldm_sd3.py b/scepter/modules/model/network/ldm/ldm_sd3.py
index c711485..b41d2bb 100644
--- a/scepter/modules/model/network/ldm/ldm_sd3.py
+++ b/scepter/modules/model/network/ldm/ldm_sd3.py
@@ -136,6 +136,7 @@ class LatentDiffusionSD3(LatentDiffusion):
@torch.no_grad()
def forward_test(self,
+ image=None,
prompt=None,
sampler='ddim',
sample_steps=20,
diff --git a/scepter/modules/model/registry.py b/scepter/modules/model/registry.py
index b1293b0..d664d74 100644
--- a/scepter/modules/model/registry.py
+++ b/scepter/modules/model/registry.py
@@ -26,6 +26,27 @@ def build_model(cfg, registry, logger=None, *args, **kwargs):
model.load_pretrained_model(pretrain_cfg)
return model
+def build_diffusion(cfg, registry, logger=None, *args, **kwargs):
+ """ After build model, load pretrained model if exists key `pretrain`.
+
+ pretrain (str, dict): Describes how to load pretrained model.
+ str, treat pretrain as model path;
+ dict: should contains key `path`, and other parameters token by function load_pretrained();
+ """
+ if not isinstance(cfg, Config):
+ raise TypeError(f'Config must be type dict, got {type(cfg)}')
+ return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
+
+def build_scheduler(cfg, registry, logger=None, *args, **kwargs):
+ if not isinstance(cfg, Config):
+ raise TypeError(f'Config must be type dict, got {type(cfg)}')
+ return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
+
+def build_diffusion_sampler(cfg, registry, logger=None, *args, **kwargs):
+ if not isinstance(cfg, Config):
+ raise TypeError(f'Config must be type dict, got {type(cfg)}')
+ return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
+
MODELS = Registry('MODELS', build_func=build_model)
TOKENIZERS = Registry('TOKENIZER', build_func=build_model)
@@ -37,3 +58,9 @@ BRICKS = Registry('BRICKS', build_func=build_model)
STEMS = BRICKS
LOSSES = Registry('LOSSES', build_func=build_model)
TUNERS = Registry('TUNERS', build_func=build_model)
+
+# reigister cls for diffusion.
+
+DIFFUSIONS = Registry('DIFFUSIONS', build_func=build_diffusion)
+NOISE_SCHEDULERS = Registry('NOISE_SCHEDULERS', build_func=build_diffusion)
+DIFFUSION_SAMPLERS = Registry('DIFFUSION_SAMPLERS', build_func=build_diffusion_sampler)
diff --git a/scepter/modules/opt/lr_schedulers/__init__.py b/scepter/modules/opt/lr_schedulers/__init__.py
index 36a5138..12b9340 100644
--- a/scepter/modules/opt/lr_schedulers/__init__.py
+++ b/scepter/modules/opt/lr_schedulers/__init__.py
@@ -3,4 +3,5 @@
from scepter.modules.opt.lr_schedulers.define_schedulers import LinoPolyLR
from scepter.modules.opt.lr_schedulers.official_schedulers import * # noqa
-from scepter.modules.opt.lr_schedulers.warmup import WarmupToConstantLR
+from scepter.modules.opt.lr_schedulers.warmup import (StepAnnealingLR,
+ WarmupToConstantLR)
diff --git a/scepter/modules/opt/lr_schedulers/warmup.py b/scepter/modules/opt/lr_schedulers/warmup.py
index f8ffe31..7a6b3e6 100644
--- a/scepter/modules/opt/lr_schedulers/warmup.py
+++ b/scepter/modules/opt/lr_schedulers/warmup.py
@@ -33,7 +33,6 @@ class WarmupToConstantLR(BaseScheduler):
WarmupToConstantLR.para_dict,
set_name=True)
-
class AnnealingLR(_LRScheduler):
def __init__(self,
optimizer,
diff --git a/scepter/modules/solver/base_solver.py b/scepter/modules/solver/base_solver.py
index 9fd18c9..5dd5734 100644
--- a/scepter/modules/solver/base_solver.py
+++ b/scepter/modules/solver/base_solver.py
@@ -18,7 +18,10 @@ from scepter.modules.solver.hooks import HOOKS
from scepter.modules.utils.config import Config, dict_to_yaml
from scepter.modules.utils.data import transfer_data_to_cuda
from scepter.modules.utils.directory import get_relative_folder, osp_path
-from scepter.modules.utils.distribute import dist, gather_data, we
+from scepter.modules.utils.distribute import (
+ dist, gather_data, we, all_reduce,
+ _serialize_to_tensor, broadcast, _unserialize_from_tensor,
+ all_reduce, barrier)
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger, init_logger
from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe,
@@ -184,6 +187,20 @@ try:
except Exception as e:
warnings.warn(f'{e}')
+def async_str(text):
+ broadcast_size = torch.zeros(1, dtype=torch.long).to(we.device_id)
+ if we.rank == 0:
+ text_tensor = _serialize_to_tensor(text).to(we.device_id)
+ broadcast_size[0] = len(text_tensor)
+ broadcast(broadcast_size, src=0)
+ broadcast(text_tensor, src=0)
+ else:
+ broadcast(broadcast_size, src=0)
+ text_tensor = torch.empty((broadcast_size[0],), dtype=torch.uint8).to(we.device_id)
+ broadcast(text_tensor, src=0)
+ text = _unserialize_from_tensor(text_tensor)
+ return text
+
class BaseSolver(object, metaclass=ABCMeta):
""" Base Solver.
@@ -198,18 +215,12 @@ class BaseSolver(object, metaclass=ABCMeta):
'description': 'The precision for train process.'
},
'FILE_SYSTEM': {},
- 'ACCU_STEP': {
- 'value':
- 1,
- 'description':
- 'When use ddp, the grad accumulate steps for each process.'
- },
'RESUME_FROM': {
'value': '',
'description': 'Resume from some state of training!'
},
'MAX_EPOCHS': {
- 'value': 10,
+ 'value': -1,
'description': 'Max epochs for training.'
},
'NUM_FOLDS': {
@@ -252,31 +263,31 @@ class BaseSolver(object, metaclass=ABCMeta):
def __init__(self, cfg, logger=None):
# initialize some hyperparameters
+ self.cfg = cfg
+ self.logger = logger
self.file_system = cfg.get('FILE_SYSTEM', None)
- self.work_dir: str = cfg.WORK_DIR
+ self.work_dir: str = async_str(cfg.WORK_DIR)
+ barrier()
self.pl_dir = self.work_dir
self.log_file = osp_path(self.work_dir, cfg.LOG_FILE)
self.optimizer, self.lr_scheduler = None, None
- self.cfg = cfg
- self.logger = logger
- self.resume_from: str = cfg.RESUME_FROM
- self.max_epochs: int = cfg.MAX_EPOCHS
+ self.resume_from: str = cfg.get("RESUME_FROM", None)
+ self.max_epochs: int = cfg.get("MAX_EPOCHS", -1)
self.use_pl = we.use_pl
self.train_precision = self.cfg.get('TRAIN_PRECISION', 32)
self._mode_set = set()
self._mode = 'train'
self.probe_ins = {}
+ self.collect_probe_ins = {}
self.clear_probe_ins = {}
self._num_folds: int = 1
if not self.use_pl:
world_size = we.world_size
if world_size > 1:
- self._num_folds: int = cfg.NUM_FOLDS
+ self._num_folds: int = cfg.get("NUM_FOLDS", 1)
if cfg.have('MODE'):
self._mode_set.add(cfg.MODE)
self._mode = cfg.MODE
- if we.is_distributed:
- self.accu_step = cfg.get('ACCU_STEP', 1)
self.do_step = True
self.hooks_dict = {'train': [], 'eval': [], 'test': []}
@@ -305,6 +316,8 @@ class BaseSolver(object, metaclass=ABCMeta):
self._prefix = FS.get_fs_client(self.work_dir).get_prefix()
if not FS.exists(self.work_dir):
FS.make_dir(self.work_dir)
+ assert self.cfg.have('MODEL')
+ self.cfg.MODEL.WORK_DIR = self.work_dir
self.logger.info(
f"Parse work dir {self.work_dir}'s prefix is {self._prefix}")
@@ -325,6 +338,7 @@ class BaseSolver(object, metaclass=ABCMeta):
def __setattr__(self, key, value):
if isinstance(value, BaseModel):
self.probe_ins[key] = value.probe_data
+ self.collect_probe_ins[key] = value.collect_probe
self.clear_probe_ins[key] = value.clear_probe
super().__setattr__(key, value)
@@ -666,7 +680,7 @@ class BaseSolver(object, metaclass=ABCMeta):
return self._iter[self._mode]
@property
- def probe_data(self):
+ def probe_data_dict(self):
return self._probe_data[self._mode]
@property
@@ -736,8 +750,32 @@ class BaseSolver(object, metaclass=ABCMeta):
else:
self._dist_data[self.mode][key][k] = v
+ @property
+ def collect_probe(self):
+ probe_data_dict = self._probe_data[self.mode]
+ for k, func in self.collect_probe_ins.items():
+ for kk, vv in func().items():
+ probe_data_dict[f'{k}/{kk}'] = vv
+ return probe_data_dict
+
@property
def probe_data(self): # noqa
+ if hasattr(self, f'{self.mode}_pre_save_paras'):
+ pre_save_paras = getattr(self, f'{self.mode}_pre_save_paras')
+ save_folder = pre_save_paras['save_folder']
+ save_probe_prefix = pre_save_paras['save_probe_prefix']
+ step = pre_save_paras['step']
+ save_image_postfix = pre_save_paras.get('save_image_postfix', 'jpg')
+ save_video_postfix = pre_save_paras.get('save_video_postfix', 'mp4')
+ for k, v in self.collect_probe.items():
+ if save_probe_prefix is not None:
+ ret_prefix = os.path.join(save_folder, save_probe_prefix)
+ else:
+ ret_prefix = os.path.join(save_folder, k.replace('/', '_') + f'_step_{step}')
+ v.presave(prefix = ret_prefix,
+ image_postfix = save_image_postfix,
+ video_postfix = save_video_postfix,
+ rank = we.rank)
gather_probe_data = gather_data(self._probe_data[self.mode])
_dist_data_list = gather_data([self._dist_data[self.mode] or {}])
if not we.rank == 0:
@@ -836,9 +874,10 @@ class BaseSolver(object, metaclass=ABCMeta):
for key in keys:
value = data_dict[key]
if isinstance(value, torch.Tensor) and value.ndim == 0:
- if dist.is_available() and dist.is_initialized():
+ if we.is_distributed:
value = value.data.clone()
- dist.all_reduce(value.div_(dist.get_world_size()))
+ all_reduce(value, group=we.data_parallel_group)
+ value = value/we.data_group_world_size
ret[key] = value
else:
ret[key] = value
@@ -911,7 +950,7 @@ class BaseSolver(object, metaclass=ABCMeta):
}
:return:
'''
- return dict_to_yaml('solvername',
+ return dict_to_yaml('SOLVER',
__class__.__name__,
BaseSolver.para_dict,
set_name=True)
diff --git a/scepter/modules/solver/diffusion_solver.py b/scepter/modules/solver/diffusion_solver.py
index 8fccecd..a4b37da 100644
--- a/scepter/modules/solver/diffusion_solver.py
+++ b/scepter/modules/solver/diffusion_solver.py
@@ -2,11 +2,14 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os
+import re
+import warnings
from collections import OrderedDict, defaultdict
+from functools import partial
import numpy as np
import torch
-import torch.cuda.amp as amp
+import torch.nn as nn
from scepter.modules.data.dataset import DATASETS
from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS
from scepter.modules.opt.optimizers import OPTIMIZERS
@@ -16,19 +19,80 @@ from scepter.modules.utils.config import Config, dict_to_yaml
from scepter.modules.utils.data import transfer_data_to_cuda
from scepter.modules.utils.distribute import we
from scepter.modules.utils.probe import ProbeData
-from torch.distributed.fsdp import (BackwardPrefetch, CPUOffload,
- FullStateDictConfig,
- FullyShardedDataParallel, MixedPrecision,
- ShardingStrategy, StateDictType)
+from torch.distributed.fsdp import FullStateDictConfig
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+from torch.distributed.fsdp import (MixedPrecision, ShardingStrategy,
+ StateDictType)
+from torch.distributed.fsdp.wrap import (lambda_auto_wrap_policy,
+ size_based_auto_wrap_policy)
from torch.nn.parallel import DistributedDataParallel
from tqdm import tqdm
sharding_strategy_map = {
'full_shard': ShardingStrategy.FULL_SHARD,
- 'shard_grad_op': ShardingStrategy.SHARD_GRAD_OP
+ 'shard_grad_op': ShardingStrategy.SHARD_GRAD_OP,
+ 'hybrid_shard': ShardingStrategy.HYBRID_SHARD
}
+def shard_model(model,
+ device_id,
+ param_dtype=torch.bfloat16,
+ reduce_dtype=torch.float32,
+ buffer_dtype=torch.float32,
+ fsdp_group = ['blocks'],
+ sharding_strategy=ShardingStrategy.FULL_SHARD,
+ sync_module_states=False):
+ wrap_modules = []
+ for module_name in fsdp_group:
+ if hasattr(model, module_name):
+ if isinstance(getattr(model, module_name), (list, tuple, nn.ModuleList)):
+ wrap_modules.extend([m for m in getattr(model, module_name)])
+ else:
+ wrap_modules.extend([getattr(model, module_name)])
+ else:
+ warnings.warn("Can't find module {} in model".format(module_name))
+ return FSDP(
+ module=model,
+ process_group=None,
+ sharding_strategy=sharding_strategy,
+ auto_wrap_policy=partial(
+ # size_based_auto_wrap_policy, min_num_params=int(1e6),
+ lambda_auto_wrap_policy,
+ lambda_fn=lambda m: m in wrap_modules),
+ mixed_precision=MixedPrecision(param_dtype=param_dtype,
+ reduce_dtype=reduce_dtype,
+ buffer_dtype=buffer_dtype),
+ device_id=device_id,
+ sync_module_states=sync_module_states)
+
+
+def get_module(instance, sub_module):
+ sub_module_list = sub_module.split('.')
+ for sub_mod in sub_module_list:
+ if sub_mod == '':
+ continue
+ if hasattr(instance, sub_mod):
+ instance = getattr(instance, sub_mod)
+ else:
+ return None
+ return instance
+
+
+def set_module(instance, sub_module, value):
+ sub_module_list = sub_module.split('.')
+ instance_list = []
+ for sub_mod in sub_module_list:
+ if hasattr(instance, sub_mod):
+ instance_list.append((instance, sub_mod))
+ instance = getattr(instance, sub_mod)
+ instance = value
+ if len(instance_list) > 0:
+ for parents_instance, sub_mod in instance_list[::-1]:
+ setattr(parents_instance, sub_mod, instance)
+ instance = parents_instance
+
+
@SOLVERS.register_class()
class LatentDiffusionSolver(BaseSolver):
para_dict = {
@@ -61,6 +125,24 @@ class LatentDiffusionSolver(BaseSolver):
'description':
f'The shard strategy for fsdp, select from {list(sharding_strategy_map.keys())}',
},
+ 'FSDP_REDUCE_DTYPE': {
+ 'value': 'float32',
+ 'description': 'The dtype of reduce in FSDP.'
+ },
+ 'FSDP_BUFFER_DTYPE': {
+ 'value': 'float32',
+ 'description': 'The dtype of buffer in FSDP.'
+ },
+ 'FSDP_SHARD_MODULES': {
+ 'value': ['model'],
+ 'description': 'The modules to be sharded in FSDP.'
+ },
+ 'SAVE_MODULES': {
+ 'value':
+ None,
+ 'description':
+ 'The modules to be saved, default is None to save all modules in checkpoint file.'
+ },
'IMAGE_LOG_STEP': {
'value': 2000,
'description': 'The interval for image log.',
@@ -112,6 +194,13 @@ class LatentDiffusionSolver(BaseSolver):
else:
self.logger.info('Use default backend.')
self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard')
+ self.reduce_dtype = getattr(torch,
+ cfg.get('FSDP_REDUCE_DTYPE', 'float32'))
+ self.buffer_dtype = getattr(torch,
+ cfg.get('FSDP_BUFFER_DTYPE', 'float32'))
+ self.shard_modules = cfg.get('FSDP_SHARD_MODULES', ['model'])
+ self.save_modules = cfg.get('SAVE_MODULES', ['model'])
+ self.train_modules = cfg.get('TRAIN_MODULES', ['model'])
self.image_log_step = cfg.get('IMAGE_LOG_STEP', 2000)
self._image_out = defaultdict(list)
self.load_model_only = cfg.get('LOAD_MODEL_ONLY', False)
@@ -120,6 +209,7 @@ class LatentDiffusionSolver(BaseSolver):
self.sample_args = cfg.get('SAMPLE_ARGS', None)
self.tuner_cfg = cfg.get('TUNER', None)
self.freeze_cfg = cfg.get('FREEZE', None)
+ self.log_train_num = cfg.get("LOG_TRAIN_NUM", -1)
def set_up(self):
self.construct_data()
@@ -178,7 +268,7 @@ class LatentDiffusionSolver(BaseSolver):
self.model = self.model.to(we.device_id)
def init_lr(self):
- rescale_lr = self.cfg.get('RESCALE_LR', True)
+ rescale_lr = self.cfg.get('RESCALE_LR', False)
if rescale_lr and 'train' in self.datas and self.cfg.have('OPTIMIZER'):
if we.world_size > 1:
all_batch_size = self.datas['train'].batch_size * we.world_size
@@ -188,32 +278,76 @@ class LatentDiffusionSolver(BaseSolver):
self.cfg.OPTIMIZER.LEARNING_RATE /= 640
def init_opti(self):
- if hasattr(self.model, 'ignored_parameters'):
- train_params, ignored_params = self.model.parameters(
- ), self.model.ignored_parameters()
- else:
- train_params, ignored_params = self.model.parameters(), None
+ import torch.cuda.amp as amp
+
if we.is_distributed:
if self.use_fairscale:
from fairscale.nn.data_parallel import ShardedDataParallel
from fairscale.optim.oss import OSS
+ if hasattr(self.model, 'ignored_parameters'):
+ train_params, ignored_params = self.model.parameters(
+ ), self.model.ignored_parameters()
+ else:
+ train_params, ignored_params = self.model.parameters(
+ ), None
self.optimizer = OSS(params=train_params,
optim=torch.optim.AdamW,
lr=self.cfg.OPTIMIZER.LEARNING_RATE)
self.model = ShardedDataParallel(self.model, self.optimizer)
elif self.use_fsdp:
- mixed_precision = MixedPrecision(param_dtype=self.dtype,
- reduce_dtype=self.dtype,
- buffer_dtype=self.dtype)
- sharding_strategy = sharding_strategy_map[self.model_shard]
- self.model = FullyShardedDataParallel(
- self.model,
- mixed_precision=mixed_precision,
- cpu_offload=CPUOffload(offload_params=False),
- sharding_strategy=sharding_strategy,
- backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
- device_id=torch.cuda.current_device(),
- ignored_parameters=ignored_params)
+ shard_fn = partial
+ if self.shard_modules is not None:
+ for module in self.shard_modules:
+ if isinstance(module, str):
+ sub_module = get_module(self.model, module)
+ if sub_module is not None:
+ sub_module = shard_model(
+ sub_module,
+ device_id=we.device_id,
+ param_dtype=self.dtype,
+ reduce_dtype=self.reduce_dtype,
+ buffer_dtype=self.buffer_dtype,
+ sharding_strategy=sharding_strategy_map[self.model_shard],
+ sync_module_states=True)
+ set_module(self.model, module, sub_module)
+ elif isinstance(module, (dict, Config)):
+ sub_module = get_module(self.model, module["MODULE"])
+ if sub_module is not None:
+ sub_module = shard_model(
+ sub_module,
+ device_id=we.device_id,
+ param_dtype=self.dtype,
+ reduce_dtype=self.reduce_dtype,
+ buffer_dtype=self.buffer_dtype,
+ fsdp_group=module.get("FSDP_GROUP", ["blocks"]),
+ sharding_strategy=sharding_strategy_map[self.model_shard],
+ sync_module_states=True)
+ set_module(self.model, module["MODULE"], sub_module)
+ else:
+ self.logger.warning(
+ 'FSDP_SHARD_MODULES is None, which means wraping the whold model as the '
+ 'fsdp instance. When using FSDP, it is necessary to specify the modules '
+ 'to be wrapped; otherwise, there may be a situation where submodules are '
+ 'not the root module, which can lead to unexpected issues. Specify the '
+ 'modules to be wrapped by setting FSDP_SHARD_MODULES to a list of modules '
+ 'that need wrapping.')
+ self.model = shard_fn(self.model)
+ train_params = []
+ if self.train_modules is None:
+ self.logger.warning(
+ 'When using FSDP, it is necessary to explicitly specify the modules to be '
+ 'trained or the modules for which gradients will be computed, otherwise, '
+ 'there will be issues with gradient calculation.')
+ assert self.train_modules is None
+ else:
+ self.logger.info(
+ f"The modules {','.join(self.train_modules)} 's parameters will be backwarded."
+ )
+ for module in self.train_modules:
+ if hasattr(self.model, module):
+ current_module = getattr(self.model, module)
+ train_params += list(current_module.parameters())
+
self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
logger=self.logger,
parameters=train_params)
@@ -223,14 +357,16 @@ class LatentDiffusionSolver(BaseSolver):
self.model,
device_ids=[torch.cuda.current_device()],
output_device=torch.cuda.current_device(),
- find_unused_parameters=True)
- self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
- logger=self.logger,
- parameters=train_params)
+ find_unused_parameters=False)
+ self.optimizer = OPTIMIZERS.build(
+ self.cfg.OPTIMIZER,
+ logger=self.logger,
+ parameters=self.model.parameters())
else:
- self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
- logger=self.logger,
- parameters=train_params)
+ self.optimizer = OPTIMIZERS.build(
+ self.cfg.OPTIMIZER,
+ logger=self.logger,
+ parameters=self.model.parameters())
if self.cfg.have('LR_SCHEDULER') and self.optimizer is not None:
self.cfg.LR_SCHEDULER.TOTAL_STEPS = self.max_steps
@@ -238,22 +374,22 @@ class LatentDiffusionSolver(BaseSolver):
logger=self.logger,
optimizer=self.optimizer)
- if self.cfg.DTYPE == 'float16':
+ if self.cfg.DTYPE in ['float16']:
if we.is_distributed:
if self.use_fairscale:
from fairscale.optim.grad_scaler import ShardedGradScaler
self.scaler = ShardedGradScaler(enabled=True)
elif self.use_fsdp:
- from torch.distributed.fsdp.sharded_grad_scaler import \
- ShardedGradScaler
- self.scaler = ShardedGradScaler()
+ from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
+ self.scaler = ShardedGradScaler(enabled=True,
+ process_group=None)
else:
self.scaler = amp.GradScaler()
else:
self.scaler = amp.GradScaler()
else:
self.scaler = None
-
+ self.logger.info(self.model)
def load_checkpoint(self, checkpoint: dict):
"""
Load checkpoint function
@@ -262,19 +398,56 @@ class LatentDiffusionSolver(BaseSolver):
"""
if 'model' in checkpoint:
if hasattr(self.model, 'module'):
- self.model.module.load_state_dict(checkpoint['model'])
+ if self.save_modules is not None:
+ for module in self.save_modules:
+ current_module = get_module(self.model.module, module)
+ if current_module is not None and module in checkpoint[
+ 'model']:
+ current_module.load_state_dict(
+ checkpoint['model'][module])
+ self.logger.info(
+ f'Load checkpoint for model.{module}')
+ else:
+ self.model.module.load_state_dict(checkpoint['model'])
+ self.logger.info('Load checkpoint for model.')
else:
- self.model.load_state_dict(checkpoint['model'])
+ if self.save_modules is not None:
+ for module in self.save_modules:
+ current_module = get_module(self.model, module)
+ self.logger.info(f'Load checkpoint for model.{module}')
+ if current_module is not None and module in checkpoint[
+ 'model']:
+ current_module.load_state_dict(
+ checkpoint['model'][module])
+ else:
+ self.model.load_state_dict(checkpoint['model'])
+ self.logger.info('Load checkpoint for model.')
else:
if hasattr(self.model, 'module'):
self.model.module.load_state_dict(checkpoint)
else:
self.model.load_state_dict(checkpoint)
+ self.logger.info('Load checkpoint for model.')
if not self.load_model_only:
if 'optimizer' in checkpoint and self.optimizer:
- self.optimizer.load_state_dict(checkpoint['optimizer'])
+ if self.use_fsdp:
+ for module in self.train_modules:
+ current_module = get_module(self.model, module)
+ if current_module is not None and module in checkpoint[
+ 'optimizer']:
+ state = FSDP.optim_state_dict_to_load(
+ current_module, self.optimizer,
+ checkpoint['optimizer'][module])
+ self.optimizer.load_state_dict(state)
+ self.logger.info(
+ f'Load checkpoint for optimizer {module}.')
+ else:
+ self.optimizer.load_state_dict(checkpoint['optimizer'])
+ self.logger.info(f'Load checkpoint for optimizer.')
if 'scaler' in checkpoint and self.scaler:
self.scaler.load_state_dict(checkpoint['scaler'])
+ self.logger.info(f'Load checkpoint for scaler.')
+ self.logger.info('Load checkpoint finished.')
def save_checkpoint(self) -> dict:
"""
@@ -282,23 +455,76 @@ class LatentDiffusionSolver(BaseSolver):
:return:
"""
ckpt = dict()
+ if not self.use_fsdp and not we.rank == 0:
+ return ckpt
+ ckpt['model'] = OrderedDict()
if we.is_distributed:
if self.use_fsdp:
save_policy = FullStateDictConfig(offload_to_cpu=True,
rank0_only=True)
- with FullyShardedDataParallel.state_dict_type(
- self.model, StateDictType.FULL_STATE_DICT,
- save_policy):
- ckpt['model'] = self.model.state_dict()
+ if self.shard_modules is not None:
+ if self.save_modules is None:
+ self.logger.warning(
+ 'When using FSDP, after specifying the modules to be wrapped '
+ 'using FSDP_SHARD_MODULES, please set the modules to be saved '
+ 'in SAVE_MODULES. If set to None, it means all modules will '
+ 'be saved. However, for nested modules, the system cannot '
+ 'determine whether they are FSDP instances and need to be '
+ 'explicitly set.')
+ assert self.save_modules is not None
+ for module in self.save_modules:
+ current_module = get_module(self.model, module)
+ if current_module is not None:
+ if isinstance(current_module, FSDP):
+ # print(module, current_module._is_root)
+ with FSDP.state_dict_type(
+ current_module,
+ StateDictType.FULL_STATE_DICT,
+ save_policy):
+ ckpt['model'][
+ module] = current_module.state_dict()
+ else:
+ ckpt['model'][
+ module] = current_module.state_dict()
+ else:
+ with FSDP.state_dict_type(self.model,
+ StateDictType.FULL_STATE_DICT,
+ save_policy):
+ ckpt['model'] = self.model.state_dict()
else:
if hasattr(self.model, 'module'):
- ckpt['model'] = self.model.module.state_dict()
+ model = self.model.module
else:
- ckpt['model'] = self.model.state_dict()
+ model = self.model
+ if self.save_modules is not None:
+ for module in self.save_modules:
+ current_module = get_module(self.model, module)
+ if current_module is not None:
+ ckpt['model'][module] = current_module.state_dict()
+ else:
+ ckpt['model'] = model.state_dict()
else:
- ckpt['model'] = self.model.state_dict()
+ if hasattr(self.model, 'module'):
+ model = self.model.module
+ else:
+ model = self.model
+ if self.save_modules is not None:
+ for module in self.save_modules:
+ current_module = get_module(self.model, module)
+ if current_module is not None:
+ ckpt['model'][module] = current_module.state_dict()
+ else:
+ ckpt['model'] = model.state_dict()
if self.optimizer and not self.use_fairscale:
- ckpt['optimizer'] = self.optimizer.state_dict()
+ if self.use_fsdp and we.is_distributed:
+ ckpt['optimizer'] = OrderedDict()
+ for module in self.train_modules:
+ if hasattr(self.model, module):
+ current_module = getattr(self.model, module)
+ ckpt['optimizer'][module] = FSDP.optim_state_dict(
+ current_module, self.optimizer)
+ else:
+ ckpt['optimizer'] = self.optimizer.state_dict()
if self.scaler:
ckpt['scaler'] = self.scaler.state_dict()
return ckpt
@@ -354,7 +580,8 @@ class LatentDiffusionSolver(BaseSolver):
})
self.current_batch_data[self.mode] = batch_data
if self.sample_args:
- batch_data.update(self.sample_args.get_lowercase_dict())
+ self.current_batch_data[self.mode].update(
+ self.sample_args.get_lowercase_dict())
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
@@ -395,58 +622,29 @@ class LatentDiffusionSolver(BaseSolver):
log_data, log_label, ori_label = [], [], []
for result in all_results:
# the inference image use
+ ret_images, ret_labels = [], []
if 'hint' in result:
- merge_image = torch.cat([
- result['hint'][:result['image'].shape[0]], result['image']
- ],
- dim=2)
- log_data.append((merge_image.permute(1, 2, 0).cpu().numpy() *
- 255).astype(np.uint8))
- else:
- log_data.append(
- (result['image'].permute(1, 2, 0).cpu().numpy() *
- 255).astype(np.uint8))
- log_label.append(result['prompt'] + ' NegPrompt: ' +
- result['n_prompt'])
+ ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
+ 255).astype(np.uint8))
+ ret_labels.append(f"Control Image")
+ ret_images.append(
+ (result['image'].permute(1, 2, 0).cpu().numpy() *
+ 255).astype(np.uint8))
+ ret_labels.append(result['prompt'] +
+ " |NegPrompt| " +
+ result['n_prompt'])
+ log_data.append(ret_images)
+ log_label.append(ret_labels)
ori_label.append(result['prompt'])
- self.register_probe({'test_label': log_label})
+ self.register_probe({'eval_label': log_label})
self.register_probe({
- 'test_image':
+ 'eval_image':
ProbeData(log_data,
is_image=True,
build_html=True,
build_label=log_label)
})
-
- log_data, log_label, ori_label = [], [], []
- for result in all_results:
- # the inference image use
- if 'train_n_image' in result:
- if 'hint' in result:
- merge_image = torch.cat([
- result['hint'][:result['train_n_image'].shape[0]],
- result['train_n_image']
- ],
- dim=2)
- log_data.append(
- (merge_image.permute(1, 2, 0).cpu().numpy() *
- 255).astype(np.uint8))
- else:
- log_data.append((result['train_n_image'].permute(
- 1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
- log_label.append(result['prompt'] + 'NegPrompt' +
- result['train_n_prompt'])
- ori_label.append(result['prompt'])
- if len(log_data) > 0:
- self.register_probe({'test_train_n_label': log_label})
- self.register_probe({
- 'test_train_n_image':
- ProbeData(log_data,
- is_image=True,
- build_html=True,
- build_label=log_label)
- })
self.after_all_iter(self.hooks_dict[self._mode])
@torch.no_grad()
@@ -468,14 +666,23 @@ class LatentDiffusionSolver(BaseSolver):
rank=we.rank)
all_results.extend(results)
self.after_iter(self.hooks_dict[self._mode])
- log_data, log_label = [], []
+ log_data, log_label, ori_label = [], [], []
for result in all_results:
# the inference image use
- log_data.append((result['image'].permute(1, 2, 0).cpu().numpy() *
- 255).astype(np.uint8))
- log_label.append(result['prompt'] +
+ ret_images, ret_labels = [], []
+ if 'hint' in result:
+ ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
+ 255).astype(np.uint8))
+ ret_labels.append(f"Control Image")
+ ret_images.append(
+ (result['image'].permute(1, 2, 0).cpu().numpy() *
+ 255).astype(np.uint8))
+ ret_labels.append(result['prompt'] +
" |NegPrompt| " +
result['n_prompt'])
+ log_data.append(ret_images)
+ log_label.append(ret_labels)
+ ori_label.append(result['prompt'])
self.register_probe({'test_label': log_label})
self.register_probe({
@@ -485,27 +692,6 @@ class LatentDiffusionSolver(BaseSolver):
build_html=True,
build_label=log_label)
})
-
- log_data, log_label = [], []
- for result in all_results:
- # the inference image use
- if 'train_n_image' in result:
- log_data.append(
- (result['train_n_image'].permute(1, 2, 0).cpu().numpy() *
- 255).astype(np.uint8))
- log_label.append(result['prompt'] +
- " |NegPrompt| " +
- result['train_n_prompt'])
- if len(log_data) > 0:
- self.register_probe({'test_train_n_label': log_label})
- self.register_probe({
- 'test_train_n_image':
- ProbeData(log_data,
- is_image=True,
- build_html=True,
- build_label=log_label)
- })
-
self.after_all_iter(self.hooks_dict[self._mode])
def add_tuner(self, tuner_cfg, model=None):
@@ -524,7 +710,7 @@ class LatentDiffusionSolver(BaseSolver):
from swift import Swift
model = Swift.prepare_model(self.model, config=swift_cfg_dict)
- # self.logger.info([(key, param.shape) for key, param in self.model.named_parameters() if param.requires_grad])
+ self.logger.info([(key, param.shape) for key, param in model.named_parameters() if param.requires_grad])
return model
def freeze(self, freeze_cfg, model=None):
@@ -582,6 +768,11 @@ class LatentDiffusionSolver(BaseSolver):
freeze_flag = sum([p in name for p in train_part]) > 0
if freeze_flag:
param.requires_grad = True
+ elif isinstance(train_part, str):
+ for name, param in freeze_model.named_parameters():
+ if re.match(train_part, name):
+ param.requires_grad = True
+ self.logger.info([(key, param.shape) for key, param in freeze_model.named_parameters() if param.requires_grad])
return model
@torch.no_grad()
@@ -602,7 +793,11 @@ class LatentDiffusionSolver(BaseSolver):
return dict_to_yaml('solvername',
__class__.__name__,
LatentDiffusionSolver.para_dict,
- set_name=True)
+ set_name=True,
+ exclude_keys=[
+ 'EXTRA_KEYS', 'TRAIN_PRECISION', 'MAX_EPOCHS',
+ 'NUM_FOLDS'
+ ])
@property
def image_out(self):
@@ -620,29 +815,32 @@ class LatentDiffusionSolver(BaseSolver):
@property
def probe_data(self):
if not we.debug and self.mode == 'train':
+ batch_data = transfer_data_to_cuda(self.current_batch_data[self.mode])
+ self.eval_mode()
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
- outputs = self.log_image(
- transfer_data_to_cuda(self.current_batch_data[self.mode]))
+ batch_data['log_num'] = self.log_train_num
+ results = self.run_step_eval(batch_data)
+ images = batch_data['image'] if 'image' in batch_data else [None] * len(results)
+ self.train_mode()
log_data, log_label = [], []
- for result in outputs:
+ for result, image in zip(results, images):
+ ret_images, ret_labels = [], []
if 'hint' in result:
- merge_image = torch.cat([
- result['orig'],
- result['hint'][:result['orig'].shape[0]],
- result['recon']
- ],
- dim=2)
- else:
- merge_image = torch.cat([result['orig'], result['recon']],
- dim=2)
- log_data.append((merge_image.permute(1, 2, 0).cpu().numpy() *
- 255).astype(np.uint8))
- log_label.append('recon image: ' + result['prompt'] +
- " |NegPrompt| " +
- result['n_prompt'])
+ ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
+ 255).astype(np.uint8))
+ if image is not None:
+ image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)
+ ret_images.append((image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
+ ret_labels.append(f'target image')
+ ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
+ ret_labels.append(result['prompt']
+ + " |NegPrompt| "
+ + result['n_prompt'])
+ log_data.append(ret_images)
+ log_label.append(ret_labels)
self.register_probe({
'train_image':
ProbeData(log_data,
@@ -651,38 +849,6 @@ class LatentDiffusionSolver(BaseSolver):
build_label=log_label)
})
self.register_probe({'train_label': log_label})
-
- # the inference image use
- log_data, log_label = [], []
- for result in outputs:
- if 'train_n_image' in result:
- if 'hint' in result:
- merge_image = torch.cat([
- result['orig'],
- result['hint'][:result['orig'].shape[0]],
- result['train_n_image']
- ],
- dim=2)
- else:
- merge_image = torch.cat(
- [result['orig'], result['train_n_image']], dim=2)
- log_data.append(
- (merge_image.permute(1, 2, 0).cpu().numpy() *
- 255).astype(np.uint8))
- log_label.append(
- 'recon image: ' + result['prompt'] +
- " |NegPrompt| " +
- result['train_n_prompt'])
-
- if len(log_data) > 0:
- self.register_probe({'train_n_label': log_label})
- self.register_probe({
- 'train_n_image':
- ProbeData(log_data,
- is_image=True,
- build_html=True,
- build_label=log_label)
- })
return super().probe_data
def print_memory_status(self):
@@ -719,12 +885,20 @@ class LatentDiffusionSolver(BaseSolver):
if logger is None:
logger = self.logger
train_param_dict = {}
- forzen_param_dict = {}
+ frozen_param_dict = {}
+ ema_param_dict = {}
all_param_numel = 0
if we.debug:
for key, _ in model.named_modules():
logger.info(f'sub modules {key}.')
for key, val in model.named_parameters():
+ if 'ema' in key:
+ sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1])
+ if sub_key in ema_param_dict:
+ ema_param_dict[sub_key] += val.numel()
+ else:
+ ema_param_dict[sub_key] = val.numel()
+ continue
if val.requires_grad:
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 2)[:2])
if sub_key in train_param_dict:
@@ -733,20 +907,26 @@ class LatentDiffusionSolver(BaseSolver):
train_param_dict[sub_key] = val.numel()
else:
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1])
- if sub_key in forzen_param_dict:
- forzen_param_dict[sub_key] += val.numel()
+ if sub_key in frozen_param_dict:
+ frozen_param_dict[sub_key] += val.numel()
else:
- forzen_param_dict[sub_key] = val.numel()
+ frozen_param_dict[sub_key] = val.numel()
all_param_numel += val.numel()
if we.debug:
logger.info(key)
train_param_numel = sum(train_param_dict.values())
- forzen_param_numel = sum(forzen_param_dict.values())
+ frozen_param_numel = sum(frozen_param_dict.values())
logger.info(
f'Load trainable params {train_param_numel} / {all_param_numel} = '
f'{train_param_numel / all_param_numel:.2%}, '
f'train part: {train_param_dict}.')
logger.info(
- f'Load forzen params {forzen_param_numel} / {all_param_numel} = '
- f'{forzen_param_numel / all_param_numel:.2%}, '
- f'forzen part: {forzen_param_dict}.')
+ f'Load frozen params {frozen_param_numel} / {all_param_numel} = '
+ f'{frozen_param_numel / all_param_numel:.2%}, '
+ f'frozen part: {frozen_param_dict}.')
+ if len(ema_param_dict) > 0:
+ ema_param_numel = sum(ema_param_dict.values())
+ logger.info(
+ f'Load ema frozen params {ema_param_numel} / {all_param_numel} = '
+ f'{ema_param_numel / all_param_numel:.2%}, '
+ f'frozen part: {ema_param_dict}.')
\ No newline at end of file
diff --git a/scepter/modules/solver/hooks/backward.py b/scepter/modules/solver/hooks/backward.py
index d9c087a..d9a7e2a 100644
--- a/scepter/modules/solver/hooks/backward.py
+++ b/scepter/modules/solver/hooks/backward.py
@@ -1,8 +1,13 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+import os
import warnings
import torch
+from scepter.modules.utils.file_system import FS
+
+from scepter.modules.utils.distribute import we
+
from scepter.modules.solver.hooks.hook import Hook
from scepter.modules.solver.hooks.registry import HOOKS
from scepter.modules.utils.config import dict_to_yaml
@@ -30,6 +35,30 @@ class BackwardHook(Hook):
'EMPTY_CACHE_STEP': {
'value': -1,
'description': 'the memory empty step!'
+ },
+ 'DO_PROFILE': {
+ 'value': False,
+ 'description': 'whether to do profiling!'
+ },
+ 'PROFILE_DIR': {
+ 'value': None,
+ 'description': 'the dir for profiling!'
+ },
+ 'PROFILE_WAIT': {
+ 'value': 1,
+ 'description': 'the wait steps for profiling!'
+ },
+ 'PROFILE_WARMUP': {
+ 'value': 1,
+ 'description': 'the warmup steps for profiling!'
+ },
+ 'PROFILE_ACTIVE': {
+ 'value': 3,
+ 'description': 'the active steps for profiling!'
+ },
+ 'REPEAT': {
+ 'value': 1,
+ 'description': 'the repeat steps for profiling!'
}
}]
@@ -40,17 +69,48 @@ class BackwardHook(Hook):
self.empty_cache_step = cfg.get('EMPTY_CACHE_STEP', -1)
self.accumulate_step = cfg.get('ACCUMULATE_STEP', 1)
self.current_step = 0
+ self.wait = cfg.get('PROFILE_WAIT', 1)
+ self.warmup = cfg.get('PROFILE_WARMUP', 1)
+ self.active = cfg.get('PROFILE_ACTIVE', 3)
+ self.repeat = cfg.get('REPEAT', 1)
+ self.do_profile = cfg.get('DO_PROFILE', False)
+ self.profile_dir = cfg.get('PROFILE_DIR', None)
+ self.profile_step = 0
+ self.prof = None
+ def before_solve(self, solver):
+ if we.rank != 0:
+ return
+ if self.profile_dir is None:
+ self.log_dir = os.path.join(solver.work_dir, 'profile')
+ self._local_log_dir, _ = FS.map_to_local(self.log_dir)
+ os.makedirs(self._local_log_dir, exist_ok=True)
+ if self.do_profile:
+ self.prof = torch.profiler.profile(
+ schedule=torch.profiler.schedule(wait=self.wait, warmup=self.warmup, active=self.active, repeat=self.repeat),
+ on_trace_ready=torch.profiler.tensorboard_trace_handler(self._local_log_dir),
+ record_shapes=True,
+ with_stack=True)
+ self.prof.start()
+ solver.logger.info(f'Profiler start ...')
+ solver.logger.info(f'Profiler: save to {self.log_dir}')
+ def profile(self, solver):
+ if self.prof is None: return
+ if we.rank == 0 and self.do_profile:
+ if self.profile_step < self.wait + self.warmup + self.active:
+ self.prof.step()
+ self.profile_step += 1
+ else:
+ self.prof.stop()
+ self.do_profile = False
+ solver.logger.info(f'Profiler stop after {self.profile_step} steps')
+ FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir)
def grad_clip(self, parameters):
torch.nn.utils.clip_grad_norm_(parameters=parameters,
max_norm=self.gradient_clip,
norm_type=2)
def after_iter(self, solver):
- if (hasattr(solver, 'use_fsdp')
- and solver.use_fsdp) and self.accumulate_step > 1:
- self.logger.info("Fsdp don't surpport gradient accumulate.")
- self.accumulate_step = 1
if solver.optimizer is not None and solver.is_train_mode:
if solver.loss is None:
warnings.warn(
@@ -58,28 +118,37 @@ class BackwardHook(Hook):
)
return
if solver.scaler is not None:
- solver.scaler.scale(solver.loss).backward()
+ solver.scaler.scale(solver.loss/self.accumulate_step).backward()
if self.gradient_clip > 0:
solver.scaler.unscale_(solver.optimizer)
self.grad_clip(solver.train_parameters())
self.current_step += 1
+ # Suppose profiler run after backward, so we need to set backward_prev_step
+ # as the previous one step before the backward step
if self.current_step % self.accumulate_step == 0:
+ self.profile(solver)
solver.scaler.step(solver.optimizer)
solver.scaler.update()
solver.optimizer.zero_grad()
else:
- solver.loss.backward()
+ (solver.loss/self.accumulate_step).backward()
if self.gradient_clip > 0:
self.grad_clip(solver.train_parameters())
self.current_step += 1
+ # Suppose profiler run after backward, so we need to set backward_prev_step
+ # as the previous one step before the backward step
if self.current_step % self.accumulate_step == 0:
+ self.profile(solver)
solver.optimizer.step()
solver.optimizer.zero_grad()
if solver.lr_scheduler:
if self.current_step % self.accumulate_step == 0:
solver.lr_scheduler.step()
if self.current_step % self.accumulate_step == 0:
+ setattr(solver, 'backward_step', True)
self.current_step = 0
+ else:
+ setattr(solver, 'backward_step', False)
solver.loss = None
if self.empty_cache_step > 0 and solver.total_iter % self.empty_cache_step == 0:
torch.cuda.empty_cache()
diff --git a/scepter/modules/solver/hooks/checkpoint.py b/scepter/modules/solver/hooks/checkpoint.py
index 451c01c..a47de45 100644
--- a/scepter/modules/solver/hooks/checkpoint.py
+++ b/scepter/modules/solver/hooks/checkpoint.py
@@ -13,7 +13,6 @@ from scepter.modules.solver.hooks.registry import HOOKS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
-from swift import push_to_hub
_DEFAULT_CHECKPOINT_PRIORITY = 300
@@ -109,35 +108,42 @@ class CheckpointHook(Hook):
or solver.total_iter == solver.max_steps - 1):
solver.logger.info(
f'Saving checkpoint after {solver.total_iter + 1} steps')
- if we.rank == 0:
- save_path = osp.join(
- solver.work_dir,
- 'checkpoints/{}-{}.pth'.format(self.save_name_prefix,
- solver.total_iter + 1))
- if not self.disable_save_snapshot:
+ save_path = osp.join(
+ solver.work_dir,
+ 'checkpoints/{}-{}.pth'.format(self.save_name_prefix,
+ solver.total_iter + 1))
+ if not self.disable_save_snapshot:
+ checkpoint = solver.save_checkpoint()
+ if we.rank == 0:
with FS.put_to(save_path) as local_path:
with open(local_path, 'wb') as f:
- checkpoint = solver.save_checkpoint()
torch.save(checkpoint, f)
+ del checkpoint
- from swift import SwiftModel
- if isinstance(solver.model, SwiftModel):
- save_path = osp.join(
- solver.work_dir,
- 'checkpoints/{}-{}'.format(self.save_name_prefix,
- solver.total_iter + 1))
+ from swift import SwiftModel
+ if isinstance(solver.model, SwiftModel) or (
+ hasattr(solver.model, 'module')
+ and isinstance(solver.model.module, SwiftModel)):
+ save_path = osp.join(
+ solver.work_dir,
+ 'checkpoints/{}-{}'.format(self.save_name_prefix,
+ solver.total_iter + 1))
+ if we.rank == 0:
local_folder, _ = FS.map_to_local(save_path)
- solver.model.save_pretrained(local_folder)
+ if hasattr(solver.model, 'module'):
+ solver.model.module.save_pretrained(local_folder)
+ else:
+ solver.model.save_pretrained(local_folder)
FS.put_dir_from_local_dir(local_folder, save_path)
- else:
- if hasattr(solver, 'save_pretrained'):
- save_path = osp.join(
- solver.work_dir,
- 'checkpoints/{}-{}'.format(self.save_name_prefix,
- solver.total_iter + 1))
- local_folder, _ = FS.map_to_local(save_path)
- FS.make_dir(local_folder)
- ckpt, cfg = solver.save_pretrained()
+ else:
+ if hasattr(solver, 'save_pretrained'):
+ save_path = osp.join(
+ solver.work_dir, 'checkpoints/{}-{}'.format(
+ self.save_name_prefix, solver.total_iter + 1))
+ local_folder, _ = FS.map_to_local(save_path)
+ FS.make_dir(local_folder)
+ ckpt, cfg = solver.save_pretrained()
+ if we.rank == 0:
with FS.put_to(
os.path.join(
local_folder,
@@ -150,6 +156,7 @@ class CheckpointHook(Hook):
'configuration.json')) as local_path:
json.dump(cfg, open(local_path, 'w'))
FS.put_dir_from_local_dir(local_folder, save_path)
+ del ckpt
if self.save_last and solver.total_iter == solver.max_steps - 1:
with FS.get_fs_client(save_path) as client:
@@ -219,8 +226,10 @@ class CheckpointHook(Hook):
with open(local_file, 'wb') as f:
torch.save(checkpoint['pre_state_dict'], f)
client.put_object_from_local_file(local_file, save_path)
+ del checkpoint
def after_all_iter(self, solver):
+ from swift import push_to_hub
if we.rank == 0:
if self.push_to_hub and self.last_ckpt:
with FS.get_dir_to_local_dir(self.last_ckpt) as local_dir:
diff --git a/scepter/modules/solver/hooks/data_probe.py b/scepter/modules/solver/hooks/data_probe.py
index b0de316..7145c4f 100644
--- a/scepter/modules/solver/hooks/data_probe.py
+++ b/scepter/modules/solver/hooks/data_probe.py
@@ -23,6 +23,26 @@ class ProbeDataHook(Hook):
'PROB_INTERVAL': {
'value': 1000,
'description': 'the interval for log print!'
+ },
+ 'SAVE_NAME_PREFIX': {
+ 'value': 'step',
+ 'description': 'the prefix for save name!'
+ },
+ 'SAVE_PROBE_PREFIX': {
+ 'value': None,
+ 'description': 'the prefix for save probe!'
+ },
+ 'SAVE_LAST': {
+ 'value': False,
+ 'description': 'whether to save last!'
+ },
+ 'SAVE_IMAGE_POSTFIX': {
+ 'value': 'jpg',
+ 'description': 'the postfix for save image!'
+ },
+ 'SAVE_VIDEO_POSTFIX': {
+ 'value': 'mp4',
+ 'description': 'the postfix for save video!'
}
}]
@@ -34,31 +54,46 @@ class ProbeDataHook(Hook):
self.save_probe_prefix = cfg.get('SAVE_PROBE_PREFIX', None)
self.save_last = cfg.get('SAVE_LAST', False)
self.save_image_postfix = cfg.get('SAVE_IMAGE_POSTFIX', 'jpg')
+ self.save_video_postfix = cfg.get('SAVE_VIDEO_POSTFIX', 'mp4')
def before_all_iter(self, solver):
- pass
-
+ if not solver.mode == 'train' and hasattr(solver, 'eval_interval'):
+ solver.eval_interval = self.prob_interval
def before_iter(self, solver):
pass
+ def get_key_level_prefix(self, key, save_folder, total_iter):
+ if self.save_probe_prefix is not None:
+ ret_prefix = os.path.join(save_folder,
+ self.save_probe_prefix)
+ else:
+ ret_prefix = os.path.join(
+ save_folder,
+ key.replace('/', '_') + f'_step_{total_iter}')
+ return ret_prefix
def after_iter(self, solver):
if solver.mode == 'train' and solver.total_iter % self.prob_interval == 0:
- probe_dict = solver.probe_data
- if we.rank == 0:
- save_folder = os.path.join(
- solver.work_dir,
- f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
+ save_folder = os.path.join(
+ solver.work_dir,
+ f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
+ )
+ setattr(solver,
+ f'{solver.mode}_pre_save_paras',
+ {
+ "save_folder":save_folder,
+ "save_probe_prefix": self.save_probe_prefix,
+ "image_postfix": self.save_image_postfix,
+ "video_postfix": self.save_video_postfix,
+ "step": solver.total_iter,
+ }
)
+ gather_probe_dict = solver.probe_data
+ if we.rank == 0:
ret_data = {}
- for k, v in probe_dict.items():
- if self.save_probe_prefix is not None:
- ret_prefix = os.path.join(save_folder,
- self.save_probe_prefix)
- else:
- ret_prefix = os.path.join(
- save_folder,
- k.replace('/', '_') + f'_step_{solver.total_iter}')
- ret_one = v.to_log(ret_prefix, self.save_image_postfix)
+ for k, v in gather_probe_dict.items():
+ ret_prefix = self.get_key_level_prefix(k, save_folder, solver.total_iter)
+ ret_one = v.to_log(ret_prefix, image_postfix=self.save_image_postfix,
+ video_postfix=self.save_video_postfix)
if (isinstance(ret_one, list)
or isinstance(ret_one, dict)) and len(ret_one) < 1:
continue
@@ -74,23 +109,29 @@ class ProbeDataHook(Hook):
def after_all_iter(self, solver):
if not solver.mode == 'train':
+ step = solver._total_iter[
+ 'train'] if 'train' in solver._total_iter else 0
+ save_folder = os.path.join(
+ solver.work_dir,
+ f'{solver.mode}_probe/{self.save_name_prefix}-{step}'
+ )
+ setattr(solver,
+ f'{solver.mode}_pre_save_paras',
+ {
+ "save_folder": save_folder,
+ "save_probe_prefix": self.save_probe_prefix,
+ "image_postfix": self.save_image_postfix,
+ "step": step,
+ "video_postfix": self.save_video_postfix
+ }
+ )
probe_dict = solver.probe_data
if we.rank == 0:
- step = solver._total_iter[
- 'train'] if 'train' in solver._total_iter else 0
- save_folder = os.path.join(
- solver.work_dir,
- f'{solver.mode}_probe/{self.save_name_prefix}-{step}')
ret_data = {}
for k, v in probe_dict.items():
- if self.save_probe_prefix is not None:
- ret_prefix = os.path.join(save_folder,
- self.save_probe_prefix)
- else:
- ret_prefix = os.path.join(
- save_folder,
- k.replace('/', '_') + f'_step_{step}')
- ret_one = v.to_log(ret_prefix, self.save_image_postfix)
+ ret_prefix = self.get_key_level_prefix(k, save_folder, step)
+ ret_one = v.to_log(ret_prefix, image_postfix=self.save_image_postfix,
+ video_postfix=self.save_video_postfix)
if (isinstance(ret_one, list)
or isinstance(ret_one, dict)) and len(ret_one) < 1:
continue
diff --git a/scepter/modules/solver/hooks/log.py b/scepter/modules/solver/hooks/log.py
index 2ef47f1..95d3d1e 100644
--- a/scepter/modules/solver/hooks/log.py
+++ b/scepter/modules/solver/hooks/log.py
@@ -124,12 +124,17 @@ class LogHook(Hook):
self.time = time.time()
self.start_time = time.time()
+ self.all_throughput = 0
self.data_time = 0
+ self.batch_size = defaultdict(dict)
def before_all_iter(self, solver):
self.time = time.time()
self.last_log_step = (solver.mode, 0)
-
+ if hasattr(solver, "datas"):
+ for k, v in solver.datas.items():
+ if hasattr(v, 'batch_size'):
+ self.batch_size[k] = v.batch_size
def before_iter(self, solver):
data_time = time.time() - self.time
self.data_time = data_time
@@ -141,16 +146,21 @@ class LogHook(Hook):
outputs = solver.iter_outputs.copy()
outputs['time'] = iter_time
outputs['data_time'] = self.data_time
- if 'batch_size' in outputs:
- batch_size = outputs.pop('batch_size')
- else:
- batch_size = 1
+ if solver.mode in self.batch_size:
+ outputs['throughput'] = int(self.batch_size[solver.mode] * we.world_size / iter_time * 86400)
+ log_agg.update(outputs, 1)
+ log_agg = log_agg.aggregate(self.log_interval)
+ if 'throughput' in log_agg:
+ log_agg['throughput'] = f"{int(log_agg['throughput'][-1])}/day"
+ if solver.mode in self.batch_size:
+ log_agg['all_throughput'] = (solver.iter + 1) * we.world_size * self.batch_size[solver.mode]
+
if self.show_gpu_mem:
- outputs['nvidia-smi'] = print_memory_status()
- log_agg.update(outputs, batch_size)
+ log_agg['nvidia-smi'] = str(print_memory_status()) +"MiB"
+
if (solver.iter + 1) % self.log_interval == 0:
_print_iter_log(solver,
- log_agg.aggregate(self.log_interval),
+ log_agg,
start_time=self.start_time,
mode=solver.mode)
self.last_log_step = (solver.mode, solver.iter + 1)
diff --git a/scepter/modules/utils/config.py b/scepter/modules/utils/config.py
index a498c74..193292a 100644
--- a/scepter/modules/utils/config.py
+++ b/scepter/modules/utils/config.py
@@ -20,7 +20,7 @@ _SECURE_KEYWORDS = [
_SECURE_VALUEWORDS = ['oss://', 'oss-'] # -> "#####"
-def dict_to_yaml(module_name, name, json_config, set_name=False):
+def dict_to_yaml(module_name, name, json_config, set_name=False, exclude_keys=[]):
'''
{ "ENV" :
{ "description" : "",
@@ -70,6 +70,9 @@ def dict_to_yaml(module_name, name, json_config, set_name=False):
yaml_str = ''
# print(level_num, json_config)
if isinstance(json_config, dict):
+ for key in exclude_keys:
+ if key in json_config:
+ json_config.pop(key)
if 'value' in json_config:
value = json_config['value']
if isinstance(value, dict):
@@ -322,39 +325,39 @@ class Config(object):
'CUDNN_DETERMINISTIC': True,
'CUDNN_BENCHMARK': False
}
- self.logger.info(
- f"ENV is not set and will use default ENV as {self.cfg_dict['ENV']}; "
- f'If want to change this value, please set them in your config.'
- )
+ # self.logger.info(
+ # f"ENV is not set and will use default ENV as {self.cfg_dict['ENV']}; "
+ # f'If want to change this value, please set them in your config.'
+ # )
else:
if 'SEED' not in self.cfg_dict['ENV']:
self.cfg_dict['ENV']['SEED'] = 2023
- self.logger.info(
- f"SEED is not set and will use default SEED as {self.cfg_dict['ENV']['SEED']}; "
- f'If want to change this value, please set it in your config.'
- )
+ # self.logger.info(
+ # f"SEED is not set and will use default SEED as {self.cfg_dict['ENV']['SEED']}; "
+ # f'If want to change this value, please set it in your config.'
+ # )
os.environ['ES_SEED'] = str(self.cfg_dict['ENV']['SEED'])
self._update_dict(self.cfg_dict)
- if load:
- self.logger.info(f'Parse cfg file as \n {self.dump()}')
+ # if load:
+ # self.logger.info(f'Parse cfg file as \n {self.dump()}')
def load_from_file(self, file_name):
- self.logger.info(f'Loading config from {file_name}')
+ # self.logger.info(f'Loading config from {file_name}')
if file_name is None or not os.path.exists(file_name):
self.logger.info(f'File {file_name} does not exist!')
self.logger.warning(
- f"Cfg file is None or doesn't exist, Skip loading config from {file_name}."
+ f"Cfg file is None or doesn't exist, Skip loading config from [{file_name}]."
)
return
if file_name.endswith('.json'):
self.cfg_dict = self._load_json(file_name)
self.logger.info(
- f'System take {file_name} as json, because we find json in this file'
+ f'Loading config from [{file_name}] as json file.'
)
elif file_name.endswith('.yaml'):
self.cfg_dict = self._load_yaml(file_name)
self.logger.info(
- f'System take {file_name} as yaml, because we find yaml in this file'
+ f'Loading config from [{file_name}] as yaml file.'
)
else:
self.logger.info(
@@ -616,6 +619,20 @@ class Config(object):
config_new[key] = val
return config_new
+ def get_uppercase_dict(self, cfg_dict=None):
+ if cfg_dict is None:
+ cfg_dict = self.get_dict()
+ config_new = {}
+ for key, val in cfg_dict.items():
+ if isinstance(key, str):
+ if isinstance(val, dict):
+ config_new[key.upper()] = self.get_uppercase_dict(val)
+ else:
+ config_new[key.upper()] = val
+ else:
+ config_new[key] = val
+ return config_new
+
@staticmethod
def get_plain_cfg(cfg=None):
if isinstance(cfg, Config):
diff --git a/scepter/modules/utils/data.py b/scepter/modules/utils/data.py
index c5e75dd..b8c811e 100644
--- a/scepter/modules/utils/data.py
+++ b/scepter/modules/utils/data.py
@@ -83,7 +83,13 @@ def transfer_data_to_cuda(data_map: dict) -> dict:
elif isinstance(value, dict):
ret[key] = transfer_data_to_cuda(value)
elif isinstance(value, (list, tuple)):
- ret[key] = type(value)([transfer_data_to_cuda(t) for t in value])
+ ret_data = []
+ for t in value:
+ if not isinstance(t, dict):
+ ret_data.append(transfer_data_to_cuda({'data': t})['data'])
+ else:
+ ret_data.append(transfer_data_to_cuda(t))
+ ret[key] = type(value)(ret_data)
else:
ret[key] = value
return ret
diff --git a/scepter/modules/utils/distribute.py b/scepter/modules/utils/distribute.py
index 19d5671..b827834 100644
--- a/scepter/modules/utils/distribute.py
+++ b/scepter/modules/utils/distribute.py
@@ -4,12 +4,16 @@ import functools
import os
import pickle
import random
+import socket
import warnings
from collections import OrderedDict
+from datetime import timedelta
import numpy as np
import torch
import torch.distributed as dist
+from torch.autograd import Function
+
from scepter.modules.utils.model import StdMsg
__all__ = [
@@ -202,6 +206,96 @@ def barrier():
dist.barrier()
+def all_gather(tensor, uniform_size=True, group=None, **kwargs):
+ world_size = dist.get_world_size(group)
+ if world_size == 1:
+ return [tensor]
+ assert tensor.is_contiguous(), \
+ 'ops.all_gather requires the tensor to be contiguous()'
+
+ if uniform_size:
+ tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
+ dist.all_gather(tensor_list, tensor, group, **kwargs)
+ return tensor_list
+ else:
+ # collect tensor shapes across GPUs
+ shape = tuple(tensor.shape)
+ shape_list = generalized_all_gather(shape, group)
+
+ # flatten the tensor
+ tensor = tensor.reshape(-1)
+ size = int(np.prod(shape))
+ size_list = [int(np.prod(u)) for u in shape_list]
+ max_size = max(size_list)
+
+ # pad to maximum size
+ if size != max_size:
+ padding = tensor.new_zeros(max_size - size)
+ tensor = torch.cat([tensor, padding], dim=0)
+
+ # all_gather
+ tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
+ dist.all_gather(tensor_list, tensor, group, **kwargs)
+
+ # reshape tensors
+ tensor_list = [
+ t[:n].view(s)
+ for t, n, s in zip(tensor_list, size_list, shape_list)
+ ]
+ return tensor_list
+
+
+def _pad_to_largest_tensor(tensor, group):
+ world_size = dist.get_world_size(group=group)
+ assert world_size >= 1, \
+ 'gather/all_gather must be called from ranks within' \
+ 'the give group!'
+ local_size = torch.tensor([tensor.numel()],
+ dtype=torch.int64,
+ device=tensor.device)
+ size_list = [
+ torch.zeros([1], dtype=torch.int64, device=tensor.device)
+ for _ in range(world_size)
+ ]
+
+ # gather tensors and compute the maximum size
+ dist.all_gather(size_list, local_size, group=group)
+ size_list = [int(size.item()) for size in size_list]
+ max_size = max(size_list)
+
+ # pad tensors to the same size
+ if local_size != max_size:
+ padding = torch.zeros((max_size - local_size, ),
+ dtype=torch.uint8,
+ device=tensor.device)
+ tensor = torch.cat((tensor, padding), dim=0)
+ return size_list, tensor
+
+
+def generalized_all_gather(data, group=None):
+ if dist.get_world_size(group) == 1:
+ return [data]
+ if group is None:
+ group = get_global_gloo_group()
+
+ tensor = _serialize_to_tensor(data, group)
+ size_list, tensor = _pad_to_largest_tensor(tensor, group)
+ max_size = max(size_list)
+
+ # receiving tensors from all ranks
+ tensor_list = [
+ torch.empty((max_size, ), dtype=torch.uint8, device=tensor.device)
+ for _ in size_list
+ ]
+ dist.all_gather(tensor_list, tensor, group=group)
+
+ data_list = []
+ for size, tensor in zip(size_list, tensor_list):
+ buffer = tensor.cpu().numpy().tobytes()[:size]
+ data_list.append(pickle.loads(buffer))
+ return data_list
+
+
@functools.lru_cache()
def get_global_gloo_group():
backend = dist.get_backend()
@@ -223,12 +317,14 @@ def reduce_scatter(output,
def all_reduce(tensor, op=dist.ReduceOp.SUM, group=None, **kwargs):
if we.is_distributed:
- return dist.all_reduce(tensor, op, group, **kwargs)
+ dist.all_reduce(tensor, op, group, **kwargs)
+ return tensor
def reduce(tensor, dst, op=dist.ReduceOp.SUM, group=None, **kwargs):
if we.is_distributed:
- return dist.reduce(tensor, dst, op, group, **kwargs)
+ dist.reduce(tensor, dst, op, group, **kwargs)
+ return tensor
def _serialize_to_tensor(data):
@@ -243,6 +339,24 @@ def _unserialize_from_tensor(recv_data):
return pickle.loads(buffer)
+def find_free_port():
+ # Copied from https://github.com/facebookresearch/detectron2/blob/main/detectron2/engine/launch.py # noqa: E501
+ sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
+ # Binding to port 0 will cause the OS to find an available port for us
+ sock.bind(('', 0))
+ port = sock.getsockname()[1]
+ sock.close()
+ # NOTE: there is still a chance the port could be taken by other processes.
+ return port
+
+
+def is_free_port(port):
+ ips = socket.gethostbyname_ex(socket.gethostname())[-1]
+ ips.append('localhost')
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
+ return all(s.connect_ex((ip, port)) != 0 for ip in ips)
+
+
def send(tensor, dst, group=None, **kwargs):
if we.is_distributed:
assert tensor.is_contiguous(
@@ -288,6 +402,144 @@ def shared_random_seed():
return all_seeds[0]
+def all_to_all(x, scatter_dim, gather_dim, group=None, **kwargs):
+ """
+ `scatter` along one dimension and `gather` along another.
+ """
+ world_size = dist.get_world_size(group) if we.is_distributed else 1
+ if world_size > 1:
+ inputs = [u.contiguous() for u in x.chunk(world_size, dim=scatter_dim)]
+ outputs = [torch.empty_like(u) for u in inputs]
+ dist.all_to_all(outputs, inputs, group=group, **kwargs)
+ x = torch.cat(outputs, dim=gather_dim).contiguous()
+ return x
+
+
+def _split(input, dim, group):
+ # skip if world_size == 1
+ rank = dist.get_rank(group=group)
+ world_size = dist.get_world_size(group=group)
+ if world_size == 1:
+ return input
+
+ # split sequence
+ assert input.size(dim) % world_size == 0
+ return input.chunk(world_size, dim=dim)[rank].contiguous()
+
+
+def _gather(input, dim, group):
+ # skip if world_size == 1
+ world_size = dist.get_world_size(group=group)
+ if world_size == 1:
+ return input
+
+ # gather sequence
+ output = all_gather(input, uniform_size=True, group=group)
+ return torch.cat(output, dim=dim).contiguous()
+
+
+class AllToAll(Function):
+ @staticmethod
+ def forward(ctx, input, scatter_dim, gather_dim, group):
+ ctx.scatter_dim = scatter_dim
+ ctx.gather_dim = gather_dim
+ ctx.group = group
+ return all_to_all(input, scatter_dim, gather_dim, group)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ return (all_to_all(grad_output, ctx.gather_dim, ctx.scatter_dim,
+ ctx.group), None, None, None)
+
+
+class GradScaler(Function):
+ @staticmethod
+ def forward(ctx, input, scale):
+ ctx.scale = scale
+ return input
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.scale != 1:
+ grad_output = grad_output * ctx.scale
+ return grad_output, None
+
+
+class AllGather(Function):
+ @staticmethod
+ def forward(ctx, input, dim, group=None):
+ ctx.dim = dim
+ ctx.group = group
+ output = all_gather(input, uniform_size=True, group=group)
+ return torch.cat(output, dim=dim).contiguous()
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ rank = dist.get_rank(group=ctx.group)
+ world_size = dist.get_world_size(group=ctx.group)
+ return grad_output.chunk(world_size,
+ dim=ctx.dim)[rank].contiguous(), None, None
+
+
+def diff_all_to_all(input, scatter_dim, gather_dim, group=None):
+ return AllToAll.apply(input, scatter_dim, gather_dim, group)
+
+
+def diff_scatter_sequence(input, dim, group=None):
+ rank = dist.get_rank(group)
+ world_size = dist.get_world_size(group)
+ output = input.chunk(world_size, dim=dim)[rank].contiguous()
+ return GradScaler.apply(output, 1. / world_size)
+
+
+def diff_gather_sequence(input, dim, group=None):
+ world_size = dist.get_world_size(group)
+ output = AllGather.apply(input, dim, group)
+ return GradScaler.apply(output, world_size)
+
+
+class SplitForwardGatherBackward(Function):
+ @staticmethod
+ def forward(ctx, input, dim, group=None, grad_scale=None):
+ ctx.dim = dim
+ ctx.group = group
+ ctx.grad_scale = grad_scale
+ return _split(input, dim, group)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.grad_scale == 'up':
+ grad_output = grad_output * dist.get_world_size(group=ctx.group)
+ elif ctx.grad_scale == 'down':
+ grad_output = grad_output / dist.get_world_size(group=ctx.group)
+ return _gather(grad_output, ctx.dim, ctx.group), None, None, None
+
+
+class GatherForwardSplitBackward(Function):
+ @staticmethod
+ def forward(ctx, input, dim, group=None, grad_scale=None):
+ ctx.dim = dim
+ ctx.group = group
+ ctx.grad_scale = grad_scale
+ return _gather(input, dim, group)
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.grad_scale == 'up':
+ grad_output = grad_output * dist.get_world_size(group=ctx.group)
+ elif ctx.grad_scale == 'down':
+ grad_output = grad_output / dist.get_world_size(group=ctx.group)
+ return _split(grad_output, ctx.dim, ctx.group), None, None, None
+
+
+def split_forward_gather_backward(input, dim, group=None, grad_scale=None):
+ return SplitForwardGatherBackward.apply(input, dim, group, grad_scale)
+
+
+def gather_forward_split_backward(input, dim, group=None, grad_scale=None):
+ return GatherForwardSplitBackward.apply(input, dim, group, grad_scale)
+
+
global we
@@ -295,7 +547,10 @@ def mp_worker(gpu, ngpus_per_node, cfg, fn, pmi_rank, world_size, work_env):
rank = pmi_rank * ngpus_per_node + gpu
work_env.device_id = gpu % ngpus_per_node
work_env.rank = rank
- dist.init_process_group(backend='nccl', world_size=world_size, rank=rank)
+ dist.init_process_group(backend='nccl',
+ world_size=world_size,
+ rank=rank,
+ timeout=timedelta(seconds=18000))
torch.backends.cudnn.deterministic = cfg.ENV.get('CUDNN_DETERMINISTIC',
True)
torch.backends.cudnn.benchmark = cfg.ENV.get('CUDNN_BENCHMARK', False)
@@ -307,11 +562,50 @@ def mp_worker(gpu, ngpus_per_node, cfg, fn, pmi_rank, world_size, work_env):
work_env.logger.info(f'PMI rank {pmi_rank}!')
work_env.logger.info(f'Nums of gpu {ngpus_per_node}!')
work_env.logger.info(
- f'Current rank {work_env.rank} current devices num {ngpus_per_node} '
+ f'Current rank {work_env.rank} current devices num {ngpus_per_node} \n'
f'current machine rank {pmi_rank} and all world size {world_size}')
+ # model parallel
+ tensor_parallel_size = cfg.ENV.get('TENSOR_PARALLEL_SIZE', 1)
+ pipeline_parallel_size = cfg.ENV.get('PIPELINE_PARALLEL_SIZE', 1)
- we.set_env(work_env.get_env())
+ env_dict = work_env.get_env()
+
+ if tensor_parallel_size * pipeline_parallel_size > 1:
+ '''
+ '''
+ assert world_size % tensor_parallel_size == 0
+ assert world_size % (tensor_parallel_size *
+ pipeline_parallel_size) == 0
+ data_parallel_size = world_size // (tensor_parallel_size *
+ pipeline_parallel_size)
+ mesh = torch.arange(world_size).view(data_parallel_size,
+ pipeline_parallel_size,
+ tensor_parallel_size)
+ index = torch.where(mesh == rank)
+ assert all(u.numel() == 1 for u in index)
+ index = [u.item() for u in index]
+ for j in range(pipeline_parallel_size):
+ for k in range(tensor_parallel_size):
+ group = dist.new_group(mesh[:, j, k].tolist())
+ if j == index[1] and k == index[2]:
+ env_dict['data_parallel_group'] = group
+ for i in range(data_parallel_size):
+ for j in range(pipeline_parallel_size):
+ group = dist.new_group(mesh[i, j, :].tolist())
+ if i == index[0] and j == index[1]:
+ env_dict['tensor_parallel_group'] = group
+ for i in range(data_parallel_size):
+ for k in range(tensor_parallel_size):
+ ranks = mesh[i, :, k].tolist()
+ group = dist.new_group(ranks)
+ if i == index[0] and k == index[2]:
+ env_dict['pipeline_parallel_group'] = group
+ env_dict['pipeline_parallel_ranks'] = ranks
+ we.set_env(env_dict)
+ work_env.logger.info(str(we))
fn(cfg)
+ torch.cuda.synchronize()
+ barrier()
class Workenv(object):
@@ -322,6 +616,7 @@ class Workenv(object):
self.rank = 0
self.world_size = 1
self.device_id = 0
+ self.backend = ''
self.device_count = 1
self.seed = 2023
self.debug = False
@@ -330,24 +625,36 @@ class Workenv(object):
self.data_online = False
self.share_storage = False
+ self.data_parallel_group = None
+ self.tensor_parallel_group = None
+ self.pipleline_parallel_group = None
+ self.pipeline_parallel_ranks = None
+
def init_env(self, config, fn, logger=None):
# if use pytorch_lightning: then direct use pytorch_lightning.
config.ENV = config.get('ENV', {})
self.seed = config.ENV.get('SEED', 2023)
self.debug = os.environ.get('ES_DEBUG', None) == 'true'
+
+ if logger is None:
+ self.logger = StdMsg(name='env')
+ else:
+ self.logger = logger
+
+ self.sys_envs = config.ENV.get('SYS_ENVS', None)
+ if self.sys_envs:
+ for k, v in self.sys_envs.items():
+ os.environ[k] = v
+ self.logger.info(f'Set env variable {k}={v}')
set_random_seed(self.seed)
- if logger is not None:
- logger.info(f'And running with seed {self.seed}!')
+ self.logger.info(f'And running with seed {self.seed}!')
+
if config.ENV.get('USE_PL', False):
self.use_pl = config.ENV.USE_PL
fn(config)
return
if hasattr(config, 'args') and hasattr(config.args, 'launcher'):
self.launcher = config.args.launcher
- if logger is None:
- self.logger = StdMsg(name='env')
- else:
- self.logger = logger
self.data_online = os.environ.get('DATA_ONLINE', None) == 'true'
self.share_storage = os.environ.get('SHARE_STORAGE', None) == 'true'
@@ -379,7 +686,8 @@ class Workenv(object):
if self.is_distributed:
self.backend = config.ENV.get('BACKEND', 'nccl')
self.sync_bn = config.ENV.get('SYNC_BN', False)
- dist.init_process_group(backend=self.backend)
+ dist.init_process_group(backend=self.backend,
+ timeout=timedelta(seconds=18000))
# dist.barrier()
self.initialized = True
if dist.is_initialized():
@@ -422,7 +730,7 @@ class Workenv(object):
self.device_count = ngpus_per_node
world_size = ngpus_per_node * pmi_world_size
self.world_size = world_size
- if self.world_size > 1:
+ if self.world_size >= 1:
self.is_distributed = True
self.initialized = True
if self.is_distributed:
@@ -434,13 +742,41 @@ class Workenv(object):
self))
def get_env(self):
- return self.__dict__
+ ret_dict = {}
+ for k, v in self.__dict__.items():
+ if isinstance(v, (list, dict, int, float, str, bool)):
+ ret_dict[k] = v
+ return ret_dict
def set_env(self, we_env):
for k, v in we_env.items():
setattr(self, k, v)
set_random_seed(self.seed)
+ def group_info(self, group):
+ group_info = f'group size: {group.size()}\n'
+ group_info += f'group rank: {group.rank()}\n'
+ group_info += f'group name: {group.name()}\n'
+ return group_info
+
+ @property
+ def data_group_world_size(self):
+ if self.data_parallel_group is not None:
+ return self.data_parallel_group.size()
+ return self.world_size
+
+ @property
+ def tensor_group_world_size(self):
+ if self.tensor_parallel_group is not None:
+ return self.tensor_parallel_group.size()
+ return 1
+
+ @property
+ def pipeline_group_world_size(self):
+ if self.pipleline_parallel_group is not None:
+ return self.pipleline_parallel_group.size()
+ return 1
+
def __str__(self):
environ_str = f'Now running in the distributed environment with world size {self.world_size}\n!'
environ_str += f'Current pod have {self.device_count} devices!\n'
@@ -448,8 +784,19 @@ class Workenv(object):
environ_str += f"Current task's global rank is {self.rank} \n"
environ_str += f"Current task's data online is set {self.data_online} \n"
environ_str += f"Current task's share storage is set {self.share_storage} \n"
- environ_str += f"Current task's global seed is set {self.seed}"
+ environ_str += f"Current task's global seed is set {self.seed} \n"
+ if self.data_parallel_group is not None:
+ environ_str += f"Current task's data parallel group: {self.group_info(self.data_parallel_group)} \n"
+ if self.pipleline_parallel_group is not None:
+ environ_str += f"Current task's pipeline parallel group: {self.group_info(self.pipleline_parallel_group)} \n"
+ if self.tensor_parallel_group is not None:
+ environ_str += f"Current task's tensor parallel group: {self.group_info(self.tensor_parallel_group)} \n"
+ environ_str += f"Current task's backend is set {self.backend} \n"
return environ_str
+ def __del__(self):
+ if we.is_distributed:
+ dist.destroy_process_group()
+
we = Workenv()
diff --git a/scepter/modules/utils/file_clients/aliyun_oss_fs.py b/scepter/modules/utils/file_clients/aliyun_oss_fs.py
index 4c21ffb..a7b87fb 100644
--- a/scepter/modules/utils/file_clients/aliyun_oss_fs.py
+++ b/scepter/modules/utils/file_clients/aliyun_oss_fs.py
@@ -298,8 +298,7 @@ class AliyunOssFs(BaseFs):
try:
_ = self._download_object_multi_part(target_path,
temp_file,
- chunk_size=50 *
- 1024 * 1024)
+ chunk_size=size // 100)
break
except Exception as e:
retry += 1
@@ -361,7 +360,7 @@ class AliyunOssFs(BaseFs):
def download_one_part(key):
while not slice_queue.empty():
- R.acquire()
+ R.acquire(timeout=60)
try:
if not slice_queue.empty():
part_number, chunk = slice_queue.get_nowait()
@@ -383,7 +382,7 @@ class AliyunOssFs(BaseFs):
target_path, chunk[0], chunk[1])
with open(temp_part_file, 'wb') as f:
f.write(data)
- R.acquire()
+ R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number':
@@ -403,7 +402,7 @@ class AliyunOssFs(BaseFs):
'Download part {} for {} error {} retry {} times!'.
format(part_number, key, e, retry))
if retry >= self._retry_times:
- R.acquire()
+ R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number': part_number,
@@ -539,7 +538,7 @@ class AliyunOssFs(BaseFs):
meta_dict[target_path] = etag
else:
local_path = None
- R.acquire()
+ R.acquire(timeout=60)
try:
data_quene.put_nowait([target_path, local_path])
except Exception:
@@ -623,7 +622,11 @@ class AliyunOssFs(BaseFs):
meta_dict = self._get_dir(target_path,
local_path=local_path,
meta_dict=copy.deepcopy(meta_dict))
- json.dump(meta_dict, open(check_file, 'w'))
+ for _ in range(5):
+ try:
+ json.dump(meta_dict, open(check_file, 'w'))
+ except:
+ time.sleep(1)
if is_tmp:
self.add_temp_file(local_path)
return local_path
@@ -698,7 +701,7 @@ class AliyunOssFs(BaseFs):
def upload_one_part(key, upload_id):
while not slice_queue.empty():
- R.acquire()
+ R.acquire(timeout=60)
try:
if not slice_queue.empty():
part_number, offset, num_to_upload = slice_queue.get_nowait(
@@ -721,7 +724,7 @@ class AliyunOssFs(BaseFs):
try:
result = _bucket.upload_part(key, upload_id,
part_number, raw_data)
- R.acquire()
+ R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number': part_number,
@@ -738,7 +741,7 @@ class AliyunOssFs(BaseFs):
'Upload part {} for {} error {} retry {} times!'.
format(part_number, key, e, retry))
if retry >= self._retry_times:
- R.acquire()
+ R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number': part_number,
@@ -965,8 +968,8 @@ class AliyunOssFs(BaseFs):
key,
lifecycle,
slash_safe=slash_safe)
- _bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ)
if set_public:
+ _bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ)
output_url = output_url.replace('%2F', '/').split('?')[0]
return output_url
except Exception as e:
@@ -1060,7 +1063,7 @@ class AliyunOssFs(BaseFs):
local_path, target_path)
else:
flg = False
- R.acquire()
+ R.acquire(timeout=60)
try:
data_quene.put_nowait([local_path, target_path, flg])
except Exception:
diff --git a/scepter/modules/utils/file_clients/base_fs.py b/scepter/modules/utils/file_clients/base_fs.py
index a7a43db..391aae7 100644
--- a/scepter/modules/utils/file_clients/base_fs.py
+++ b/scepter/modules/utils/file_clients/base_fs.py
@@ -39,6 +39,7 @@ class BaseFs(object, metaclass=ABCMeta):
self._temp_files = set()
self.cfg = cfg
self.tmp_dir = cfg.get('TEMP_DIR', None)
+ self.enable_md5_path = cfg.get('ENABLE_MD5_PATH', True)
self.auto_clean = cfg.get('AUTO_CLEAN', False)
if self.tmp_dir is None:
self.auto_clean = True
@@ -306,7 +307,10 @@ class BaseFs(object, metaclass=ABCMeta):
rand_name += f'{suffix}'
tmp_file = osp.join(tempfile.gettempdir(), rand_name)
else:
- cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix)
+ if self.enable_md5_path:
+ cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix)
+ else:
+ cache_name = ''
tmp_file = osp.join(self.tmp_dir, cache_name)
return tmp_file
diff --git a/scepter/modules/utils/file_clients/huggingface_fs.py b/scepter/modules/utils/file_clients/huggingface_fs.py
index d51dc4d..1677591 100644
--- a/scepter/modules/utils/file_clients/huggingface_fs.py
+++ b/scepter/modules/utils/file_clients/huggingface_fs.py
@@ -84,6 +84,7 @@ class HuggingfaceFs(BaseFs):
target_path,
local_path=None,
wait_finish=False,
+ multi_thread=False,
timeout=3600,
sign_key=None,
worker_id=-1) -> Optional[str]:
diff --git a/scepter/modules/utils/file_clients/local_fs.py b/scepter/modules/utils/file_clients/local_fs.py
index 7be8a6e..acf7a7c 100644
--- a/scepter/modules/utils/file_clients/local_fs.py
+++ b/scepter/modules/utils/file_clients/local_fs.py
@@ -41,9 +41,7 @@ class LocalFs(BaseFs):
if target_path.startswith(self.get_prefix()):
return target_path
if target_path.startswith('./') or target_path.startswith('../'):
- return os.path.join(self.get_prefix(),
- target_path).replace('/./',
- '/').replace('/../', '/')
+ return os.path.abspath(os.path.join(self.get_prefix(), target_path))
if target_path.startswith('/'):
return target_path
if target_path.startswith('file://'):
@@ -236,8 +234,17 @@ class LocalFs(BaseFs):
def put_object(self, local_data, target_path) -> bool:
target_path = self.reconstruct_path(target_path)
- with open(target_path, 'w') as f:
- f.write(local_data)
+ dirname = os.path.dirname(target_path)
+ if not os.path.exists(dirname):
+ os.makedirs(dirname, exist_ok=True)
+ if isinstance(local_data, str):
+ with open(target_path, 'w') as f:
+ f.write(local_data)
+ elif isinstance(local_data, bytes):
+ with open(target_path, 'wb') as f:
+ f.write(local_data)
+ else:
+ raise NotImplementedError
return True
def walk_dir(self, file_dir, recurse=True):
@@ -249,6 +256,9 @@ class LocalFs(BaseFs):
def put_object_from_local_file(self, local_path, target_path) -> bool:
target_path = self.reconstruct_path(target_path)
local_path = self.reconstruct_path(local_path)
+ dirname = os.path.dirname(target_path)
+ if not os.path.exists(dirname):
+ os.makedirs(dirname, exist_ok=True)
if local_path != target_path:
try:
shutil.copy(local_path, target_path)
@@ -271,7 +281,7 @@ class LocalFs(BaseFs):
return False
return True
try:
- os.makedirs(target_dir)
+ os.makedirs(target_dir, exist_ok=True)
except Exception as e:
self.logger.error(e)
return False
@@ -309,6 +319,9 @@ class LocalFs(BaseFs):
multi_thread=False) -> bool:
local_dir = self.reconstruct_path(local_dir)
target_dir = self.reconstruct_path(target_dir)
+ dirname = os.path.dirname(target_dir)
+ if not os.path.exists(dirname):
+ os.makedirs(dirname, exist_ok=True)
if local_dir == target_dir:
return True
# # cp -f local_dir/* target_dir/*
diff --git a/scepter/modules/utils/file_clients/modelscope_fs.py b/scepter/modules/utils/file_clients/modelscope_fs.py
index bf4fb26..e3068a4 100644
--- a/scepter/modules/utils/file_clients/modelscope_fs.py
+++ b/scepter/modules/utils/file_clients/modelscope_fs.py
@@ -24,8 +24,8 @@ class ModelscopeFs(BaseFs):
super(ModelscopeFs, self).__init__(cfg, logger=logger)
retry_times = cfg.get('RETRY_TIMES', 10)
self._retry_times = retry_times
- self._model_id_loaded = set()
- self._model_file_loaded = set()
+ self._model_id_loaded = dict()
+ self._model_file_loaded = dict()
def get_prefix(self) -> str:
return 'ms://'
@@ -69,14 +69,14 @@ class ModelscopeFs(BaseFs):
if revision is not None:
local_path = local_path + '_' + str(revision)
+ model_file = os.path.join(key, file_path)
retry = 0
while retry < self._retry_times:
try:
- model_file = os.path.join(key, file_path)
if model_file in self._model_file_loaded:
- local_path = os.path.join(local_path, model_file)
+ local_path = self._model_file_loaded[model_file]
if not osp.exists(local_path):
- self._model_file_loaded.remove(key)
+ self._model_file_loaded.pop(model_file)
else:
if token is not None:
cookies = self.get_modelscope_cookie(token)
@@ -94,7 +94,7 @@ class ModelscopeFs(BaseFs):
if retry >= self._retry_times:
return None
- self._model_file_loaded.add(model_file)
+ self._model_file_loaded[model_file] = local_path
if is_tmp:
self.add_temp_file(local_path)
return local_path
@@ -142,9 +142,9 @@ class ModelscopeFs(BaseFs):
while retry < self._retry_times:
try:
if key in self._model_id_loaded:
- local_path = os.path.join(local_path, key)
+ local_path = self._model_id_loaded[key]
if not osp.exists(local_path):
- self._model_id_loaded.remove(key)
+ self._model_id_loaded.pop(key)
else:
if token is not None:
cookies = self.get_modelscope_cookie(token)
@@ -162,7 +162,7 @@ class ModelscopeFs(BaseFs):
if retry >= self._retry_times:
return None
- self._model_id_loaded.add(key)
+ self._model_id_loaded[key] = local_path
if is_tmp:
self.add_temp_file(local_path)
if not ret_folder == '':
diff --git a/scepter/modules/utils/file_system.py b/scepter/modules/utils/file_system.py
index 6fc2313..cb1a775 100644
--- a/scepter/modules/utils/file_system.py
+++ b/scepter/modules/utils/file_system.py
@@ -281,7 +281,7 @@ class FileSystem(object):
else:
return False
- def get_batch_objects_from(self, target_path_list, wait_finish=False):
+ def get_batch_objects_from(self, target_path_list, wait_finish=False, return_target_path=False):
data_quene = Queue()
batch_size = 20
R = threading.Lock()
@@ -293,7 +293,7 @@ class FileSystem(object):
wait_finish=wait_finish)
else:
local_path = None
- R.acquire()
+ R.acquire(timeout = 2)
try:
data_quene.put_nowait([target_path, local_path])
except Exception:
@@ -320,7 +320,10 @@ class FileSystem(object):
for target_path in batch_list:
local_path = file_dict.get(target_path, None)
- yield local_path
+ if return_target_path:
+ yield target_path, local_path
+ else:
+ yield local_path
def put_batch_objects_to(self,
local_path_list,
@@ -352,7 +355,7 @@ class FileSystem(object):
pass
else:
flg = False
- R.acquire()
+ R.acquire(timeout=2)
try:
data_quene.put_nowait([local_path, target_path, flg])
except Exception:
diff --git a/scepter/modules/utils/index.py b/scepter/modules/utils/index.py
index f70e791..d90eb58 100644
--- a/scepter/modules/utils/index.py
+++ b/scepter/modules/utils/index.py
@@ -6,7 +6,6 @@ from tqdm import tqdm
from scepter.modules.utils.file_system import FS
-
def init_1level_llfs(list_file,
max_lines=1024,
index_name='index',
@@ -14,6 +13,57 @@ def init_1level_llfs(list_file,
r"""Construct large-list-file index.
"""
index_dir = osp.splitext(list_file)[0]
+ total_failed = 0
+ with FS.get_from(list_file, wait_finish=True) as local_path:
+ _num_split, index_stack, save_files = 0, [], []
+ cache_data_list, target_path_list = [], []
+ with open(local_path, 'r', buffering=1000000) as f:
+ for line in tqdm(f):
+ index_stack.append(line.strip())
+ if len(index_stack) >= max_lines:
+ save_file = f'{index_dir}/{index_name}/{_num_split + 1:09d}.txt'
+ cache_data_list.append('\n'.join(index_stack).encode())
+ target_path_list.append(save_file)
+ # with FS.put_to(save_file) as cache_path:
+ # with open(cache_path, 'w') as f_w:
+ # f_w.write('\n'.join(index_stack))
+ if len(target_path_list) >= 1000:
+ put_res = [(target_path, flg) for local_path, target_path, flg
+ in FS.put_batch_objects_to(cache_data_list, target_path_list, batch_size=40)]
+ for put_flg in put_res:
+ if not put_flg:
+ total_failed += 1
+ cache_data_list, target_path_list = [], []
+ index_stack = []
+ _num_split += 1
+ save_files.append(save_file)
+ put_res = [(target_path, flg) for local_path, target_path, flg
+ in FS.put_batch_objects_to(cache_data_list, target_path_list, batch_size=50)]
+ for put_flg in put_res:
+ if not put_flg:
+ total_failed += 1
+ print(f'Failed to put {total_failed} files.')
+ if len(index_stack) > 0:
+ save_file = f'{index_dir}/{index_name}/{_num_split + 1:06d}.txt'
+ with FS.put_to(save_file) as cache_path:
+ with open(cache_path, 'w') as f_w:
+ f_w.write('\n'.join(index_stack))
+ save_files.append(save_file)
+
+ # output meta-file
+ index_file = osp.join(index_dir, f'{index_name}.txt')
+ with FS.put_to(index_file) as cache_path:
+ with open(cache_path, 'w') as f_w:
+ f_w.write('\n'.join(save_files))
+ return index_file
+
+def init_1level_llfs_single_threading(list_file,
+ max_lines=1024,
+ index_name='index',
+ delimiter='\n'):
+ r"""Construct large-list-file index.
+ """
+ index_dir = osp.splitext(list_file)[0]
print(list_file)
with FS.get_from(list_file, wait_finish=True) as local_path:
print(local_path)
diff --git a/scepter/modules/utils/logger.py b/scepter/modules/utils/logger.py
index caf51b3..99c4e97 100644
--- a/scepter/modules/utils/logger.py
+++ b/scepter/modules/utils/logger.py
@@ -47,7 +47,7 @@ def time_since(since, percent):
return '{} {:.2f}%({})'.format(as_time(s), 100 * percent, as_time(rs))
-def get_logger(name='torch dist'):
+def get_logger(name='scepter', level=logging.INFO):
logger = logging.getLogger(name)
logger.propagate = False
if len(logger.handlers) == 0:
@@ -57,8 +57,8 @@ def get_logger(name='torch dist'):
'[File: %(filename)s Function: %(funcName)s at line %(lineno)d] %(message)s'
)
std_handler.setFormatter(formatter)
- std_handler.setLevel(logging.INFO)
- logger.setLevel(logging.INFO)
+ std_handler.setLevel(level)
+ logger.setLevel(level)
logger.addHandler(std_handler)
return logger
diff --git a/scepter/modules/utils/probe.py b/scepter/modules/utils/probe.py
index cbdefe2..ab9f9d9 100644
--- a/scepter/modules/utils/probe.py
+++ b/scepter/modules/utils/probe.py
@@ -1,7 +1,9 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
+import json
import os.path
+from io import BytesIO
from numbers import Number
import numpy as np
@@ -76,7 +78,8 @@ def merge_gathered_probe(all_gathered_data):
is_image=ret_data.is_image,
build_html=ret_data.build_html,
build_label=ret_data.build_label,
- view_distribute=ret_data.view_distribute)
+ view_distribute=ret_data.view_distribute,
+ is_presave=ret_data.is_presave)
elif isinstance(ret_data.data, list):
if ret_data.build_label is not None:
if isinstance(ret_data.build_label, str):
@@ -95,19 +98,57 @@ def merge_gathered_probe(all_gathered_data):
is_image=ret_data.is_image,
build_html=ret_data.build_html,
build_label=ret_data.build_label,
- view_distribute=ret_data.view_distribute)
+ view_distribute=ret_data.view_distribute,
+ is_presave=ret_data.is_presave)
else:
all_gathered_data[key] = gathered_data
return all_gathered_data
+class MediaHandler():
+ def __init__(self, batch_size = 10):
+ self.file_list = []
+ self.target_path_list = []
+ self.target_status = {}
+ self.batch_size = batch_size
+
+ def append(self, source_file, target_path):
+ self.file_list.append(source_file)
+ self.target_path_list.append(target_path)
+ if len(self.file_list) > 2 * self.batch_size:
+ generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
+ for local_path, target_path, flg in generator:
+ self.target_status[target_path] = flg
+ self.file_list.clear()
+ self.target_path_list.clear()
+ def sync(self):
+ if len(self.file_list) > 0:
+ if len(self.file_list) > 4 * self.batch_size:
+ generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
+ for local_path, target_path, flg in generator:
+ self.target_status[target_path] = flg
+ else:
+ for file_, target_path in zip(self.file_list, self.target_path_list):
+ self.target_status[target_path] = FS.put_object(file_.getvalue(), target_path)
+ self.file_list.clear()
+ self.target_path_list.clear()
+
+ def clear(self):
+ self.file_list.clear()
+ self.target_path_list.clear()
+ self.target_status.clear()
+
+
class ProbeData():
def __init__(self,
data,
is_image=False,
+ is_video=False,
+ fps=8,
build_html=False,
build_label=None,
- view_distribute=False):
+ view_distribute=False,
+ is_presave = False):
''' Probe Data Initialize.
We only support basic types such as [torch.Tensor, numpy.ndarray, number, str],
or [dict, list] of [dict, list,
@@ -164,6 +205,28 @@ class ProbeData():
elif isinstance(v, np.ndarray):
data[k] = v
self.basic_type = False
+ elif isinstance(v, list):
+ for v_idx, v_v in enumerate(v):
+ if not check_legal_type(v_v):
+ if isinstance(v_v, torch.Tensor):
+ data[k][v_idx] = v_v.detach().cpu().numpy()
+ self.basic_type = False
+ elif isinstance(v_v, np.ndarray):
+ data[k][v_idx] = v_v
+ self.basic_type = False
+ else:
+ raise f'Unsupport data type for {v_v}'
+ elif isinstance(v, dict):
+ for k_k, v_v in v.items():
+ if not check_legal_type(v_v):
+ if isinstance(v_v, torch.Tensor):
+ data[k][k_k] = v_v.detach().cpu().numpy()
+ self.basic_type = False
+ elif isinstance(v_v, np.ndarray):
+ data[k][k_k] = v_v
+ self.basic_type = False
+ else:
+ raise f'Unsupport data type for {v_v}'
else:
raise f'Unsupport data type for {v}'
self.data = data
@@ -176,6 +239,28 @@ class ProbeData():
elif isinstance(v, np.ndarray):
data[idx] = v
self.basic_type = False
+ elif isinstance(v, list):
+ for v_idx, v_v in enumerate(v):
+ if not check_legal_type(v_v):
+ if isinstance(v_v, torch.Tensor):
+ data[idx][v_idx] = v_v.detach().cpu().numpy()
+ self.basic_type = False
+ elif isinstance(v_v, np.ndarray):
+ data[idx][v_idx] = v_v
+ self.basic_type = False
+ else:
+ raise f'Unsupport data type for {v_v}'
+ elif isinstance(v, dict):
+ for k, v_v in v.items():
+ if not check_legal_type(v_v):
+ if isinstance(v_v, torch.Tensor):
+ data[idx][k] = v_v.detach().cpu().numpy()
+ self.basic_type = False
+ elif isinstance(v_v, np.ndarray):
+ data[idx][k] = v_v
+ self.basic_type = False
+ else:
+ raise f'Unsupport data type for {v_v}'
else:
raise f'Unsupport data type for {v}'
self.data = data
@@ -185,7 +270,13 @@ class ProbeData():
raise f'Unsupport data type for {data}'
self.is_image = is_image
+ self.is_video = is_video
+ self.image_postfix = 'jpg'
+ self.video_postfix = 'mp4'
+ self.fps = fps
self.build_html = build_html
+ self.media_handler = MediaHandler()
+ self.is_presave = is_presave
if self.build_html:
assert build_label is not None
@@ -199,7 +290,84 @@ class ProbeData():
build_label, dict)
self.build_label = build_label
- def save_image(self, file_prefix, images, image_postfix):
+ def get_format(self, extension):
+ if extension.lower() in ['jpg', 'jpeg']:
+ return 'JPEG'
+ if extension.lower() in ['png']:
+ return 'PNG'
+ return 'JPEG'
+ def save_one_video(self, file_path, videos, fps = 8):
+ # write video
+ import imageio
+ try:
+ writer = imageio.get_writer(file_path, fps=fps, format=".mp4", codec='libx264', quality=8)
+ for frame in videos:
+ writer.append_data(frame)
+ writer.close()
+ return True
+ except:
+ return False
+
+
+ def save_video(self, file_prefix, videos, video_postfix, fps = 8, rank = 0):
+ if isinstance(videos, list):
+ for video in videos:
+ if isinstance(video, list):
+ raise f"Only surpport one layer nested list."
+ return [self.save_video(file_prefix + f'_{rank}_{idx}', v, video_postfix, fps) for idx, v in enumerate(videos)]
+ np_shape = videos.shape
+ # 4D
+ shape_str = '_'.join([str(v) for v in np_shape])
+ if len(np_shape) == 5:
+ # channel is 1 or 3
+ if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
+ if np_shape[-1] == 1:
+ videos = videos.reshape(videos.shape[:-1])
+ file_list = []
+ for idx in range(np_shape[0]):
+ if videos[idx].shape[0] > 1:
+ file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{video_postfix}')
+ byio = BytesIO()
+ is_suc = self.save_one_video(byio, videos[idx], fps)
+ if not is_suc:
+ byio.write(b"")
+ else:
+ file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{self.image_postfix}')
+ byio = BytesIO()
+ Image.fromarray(videos[idx][0]).save(byio, self.get_format(self.image_postfix))
+ self.media_handler.append(byio, file_path)
+ file_list.append(file_path)
+ return file_list
+ else:
+ raise f"Ensure your data's dim is BWHC, and channel is 1 or 3 for {file_prefix}"
+ elif len(np_shape) == 4:
+ if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
+ if np_shape[-1] == 1:
+ videos = videos.reshape(videos.shape[:-1])
+ if videos.shape[0] > 1:
+ file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{video_postfix}'
+ byio = BytesIO()
+ is_suc = self.save_one_video(byio, videos, fps)
+ if not is_suc:
+ byio.write(b"")
+ else:
+ file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{self.image_postfix}'
+ byio = BytesIO()
+ Image.fromarray(videos[0]).save(byio, self.get_format(self.image_postfix))
+ self.media_handler.append(byio, file_path)
+ return file_path
+ else:
+ videos = videos.reshape(list(videos.shape) + [1])
+ return self.save_video(file_prefix, videos, video_postfix, fps = fps)
+ else:
+ raise f"Ensure your data's dim is BFWHC or FWHC, and channel is 1 or 3 for {file_prefix}"
+
+ def save_image(self, file_prefix, images, image_postfix, rank = 0):
+ if isinstance(images, list):
+ for image in images:
+ if isinstance(image, list):
+ raise f"Only surpport one layer nested list."
+ return [self.save_image(file_prefix + f'_{rank}_{idx}', v, image_postfix) for idx, v in enumerate(images)]
np_shape = images.shape
# 4D
shape_str = '_'.join([str(v) for v in np_shape])
@@ -210,9 +378,10 @@ class ProbeData():
images = images.reshape(images.shape[:-1])
file_list = []
for idx in range(np_shape[0]):
- file_path = file_prefix + f'_probe_{idx}_[{shape_str}].{image_postfix}'
- with FS.put_to(file_path) as local_path:
- Image.fromarray(images[idx, ...]).save(local_path)
+ file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{image_postfix}')
+ byio = BytesIO()
+ Image.fromarray(images[idx, ...]).save(byio, self.get_format(self.image_postfix))
+ self.media_handler.append(byio, file_path)
file_list.append(file_path)
return file_list
else:
@@ -221,26 +390,29 @@ class ProbeData():
if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
if np_shape[-1] == 1:
images = images.reshape(images.shape[:-1])
- file_path = file_prefix + f'_probe_[{shape_str}].{image_postfix}'
- with FS.put_to(file_path) as local_path:
- Image.fromarray(images).save(local_path)
+ file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
+ byio = BytesIO()
+ Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
+ self.media_handler.append(byio, file_path)
return file_path
else:
images = images.reshape(list(images.shape) + [1])
return self.save_image(file_prefix, images, image_postfix)
elif len(np_shape) == 2:
- file_path = file_prefix + f'_probe_[{shape_str}].{image_postfix}'
- with FS.put_to(file_path) as local_path:
- Image.fromarray(images).save(local_path)
+ file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
+ byio = BytesIO()
+ Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
+ self.media_handler.append(byio, file_path)
return file_path
else:
raise f"Ensure your data's dim is BWHC or WHC or WH, and channel is 1 or 3 for {file_prefix}"
- def save_npy(self, file_prefix, data):
+ def save_npy(self, file_prefix, data, rank = 0):
shape_str = '_'.join([str(v) for v in data.shape])
- file_path = file_prefix + f'_{shape_str}.npy'
- with FS.put_to(file_path) as local_path:
- np.save(local_path, data)
+ file_path = file_prefix + f'_{rank}_{shape_str}.npy'
+ byio = BytesIO()
+ np.save(byio, data)
+ self.media_handler.append(byio, file_path)
return file_path
def save_html(self, html_prefix, ret_data, ret_label):
@@ -249,9 +421,10 @@ class ProbeData():
with open(local_path, 'w') as f:
f.writelines('\n')
f.writelines('\n')
+ 'opacity:1.0;} textarea {font-size: 32px;}\n')
f.writelines('
\n')
all_ranks = list()
+ is_textarea = False
for save_id, save_data in enumerate(zip(ret_data, ret_label)):
save_path, save_label = save_data
one_rank = ''
@@ -259,15 +432,32 @@ class ProbeData():
one_path, one_label = one_data
one_label = one_label.replace('<', '<').replace(
'>', '>')
- url = FS.get_url(one_path,
- lifecycle=3600 * 365 * 24).replace(
- '.oss-internal.aliyun-inc.',
- '.oss.aliyuncs.').replace(
- '-internal', '')
- one_rank += (
- f''
- f' {save_id}-{idx}|{one_label}
| '
- )
+ try:
+ url = FS.get_url(one_path,
+ lifecycle=3600 * 365 * 24).replace(
+ '.oss-internal.aliyun-inc.',
+ '.oss.aliyuncs.').replace(
+ '-internal', '')
+ except:
+ url = one_path
+ if len(one_label) > 10 and idx == len(save_path) - 1:
+ is_textarea = True
+ if self.is_video and one_path.endswith(self.video_postfix):
+ one_rank += f''
+ if idx == len(save_path) - 1 and is_textarea:
+ one_rank += f' {save_id}-{idx}
| '
+ one_rank += f' {save_id}-{idx}
| '
+ else:
+ one_rank += f'
{save_id}-{idx}|{one_label}
'
+ else:
+ one_rank += f''
+ if idx == len(save_path) - 1 and is_textarea:
+ one_rank += f' {save_id}-{idx}
| '
+ one_rank += f' {save_id}-{idx}
| '
+ else:
+ one_rank += f'
{save_id}-{idx}|{one_label}
'
+
one_rank += '
'
all_ranks.append(one_rank)
f.writelines('\n'.join(all_ranks))
@@ -277,117 +467,132 @@ class ProbeData():
def distribute(self):
return self._distribute_dict
- def to_log(self, prefix=None, image_postfix='jpg'):
+ def save_one_media(self, idx, v, prefix_path, image_postfix, video_postfix, rank = 0):
+ ret_label = None
+ if self.is_image:
+ ret_medias = self.save_image(prefix_path, v,
+ image_postfix, rank = rank)
+ elif self.is_video:
+ ret_medias = self.save_video(prefix_path, v,
+ video_postfix, fps=self.fps, rank = rank)
+ else:
+ ret_data = self.save_npy(prefix_path, v, rank = rank)
+ return ret_data, ret_label
+ ret_data = ret_medias if isinstance(ret_medias, list) else [ret_medias]
+ if self.build_html:
+ if isinstance(ret_medias, list):
+ if isinstance(self.build_label, str):
+ ret_label = [self.build_label for _ in ret_medias]
+ elif isinstance(self.build_label[idx], list):
+ assert len(self.build_label[idx]) == len(
+ ret_medias)
+ ret_label = self.build_label[idx]
+ else:
+ ret_label = [
+ self.build_label[idx]
+ for _ in ret_medias
+ ]
+ else:
+ if isinstance(self.build_label, str):
+ ret_label = [self.build_label]
+ else:
+ ret_label = [self.build_label[idx]]
+ return ret_data, ret_label
+
+ def presave(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
+ self.image_postfix = image_postfix
+ self.video_postfix = video_postfix
if isinstance(self.data, np.ndarray):
if prefix is None:
raise 'You should provide the save prefix for array sample.'
# save jpg
if self.is_image:
- ret_data = self.save_image(prefix, self.data, image_postfix)
- if isinstance(ret_data, list):
- ret_data = [ret_data]
- if self.build_html:
- ret_label = []
- if isinstance(self.build_label, str):
- ret_label.append(
- [self.build_label for _ in ret_data[0]])
- else:
- ret_label.append(self.build_label)
- if not len(ret_data[0]) == len(ret_label[0]):
- raise f"The {prefix} label's length should be equal with 1st dim {self.data.shape[0]}."
- html_prefix = prefix + '_probe.html'
- html_file = self.save_html(html_prefix, ret_data,
- ret_label)
- return {'ori_file': ret_data, 'html': html_file}
- else:
- return {'ori_file': ret_data}
- else:
- return ret_data
+ ret_data = self.save_image(prefix, self.data, image_postfix, rank = rank)
+ elif self.is_video:
+ ret_data = self.save_video(prefix, self.data, video_postfix, fps=self.fps, rank = rank)
else:
- ret_data = self.save_npy(prefix, self.data)
- return ret_data
+ ret_data = self.save_npy(prefix, self.data, rank = rank)
+ self.media_handler.sync()
+ self.media_handler.clear()
+ if isinstance(ret_data, list):
+ ret_data = [ret_data]
+ ret_label = []
+ if self.build_html:
+ if isinstance(self.build_label, str):
+ ret_label.append(
+ [self.build_label for _ in ret_data[0]])
+ else:
+ ret_label.append(self.build_label)
+ if not len(ret_data[0]) == len(ret_label[0]):
+ raise f"The {prefix} label's length should be equal with 1st dim {self.data.shape[0]}."
+ self.data = {"ret_data": ret_data, "ret_label": ret_label}
+ else:
+ self.data = ret_data
+ self.is_presave = True
elif isinstance(self.data, list):
if not self.basic_type:
ret_data = []
ret_label = []
for idx, v in enumerate(self.data):
prefix_path = os.path.join(prefix, f'{idx}')
- if self.is_image:
- ret_images = self.save_image(prefix_path, v,
- image_postfix)
- ret_data.append(ret_images if isinstance(
- ret_images, list) else [ret_images])
- if self.build_html:
- if isinstance(ret_images, list):
- if isinstance(self.build_label, str):
- ret_label.append(
- [self.build_label for _ in ret_images])
- elif isinstance(self.build_label[idx], list):
- assert len(self.build_label[idx]) == len(
- ret_images)
- ret_label.append(self.build_label[idx])
- else:
- ret_label.append([
- self.build_label[idx]
- for _ in ret_images
- ])
- else:
- if isinstance(self.build_label, str):
- ret_label.append([self.build_label])
- else:
- ret_label.append([self.build_label[idx]])
- else:
- ret_data.append(self.save_npy(prefix_path, v))
- if self.is_image and self.build_html:
- html_prefix = prefix + '_probe.html'
- html_file = self.save_html(html_prefix, ret_data,
- ret_label)
- return {'ori_file': ret_data, 'html': html_file}
- else:
- return {'ori_file': ret_data}
- else:
- return self.data
+ ret_one_data, ret_one_label = self.save_one_media(idx, v,
+ prefix_path,
+ image_postfix,
+ video_postfix,
+ rank=rank)
+ ret_data.append(ret_one_data)
+ ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
+ self.media_handler.sync()
+ self.media_handler.clear()
+ self.data = {"ret_data": ret_data, "ret_label": ret_label}
+ self.is_presave = True
elif isinstance(self.data, dict):
if not self.basic_type:
ret_data = []
ret_label = []
- for k, v in self.data:
+ for k, v in self.data.items():
prefix_path = os.path.join(prefix, f'{k}_')
- if self.is_image:
- ret_images = self.save_image(prefix_path, v,
- image_postfix)
- if isinstance(ret_images, list):
- ret_data.append(ret_images)
- else:
- ret_data.append([ret_images])
- if self.build_html:
- if isinstance(ret_images, list):
- if isinstance(self.build_label, str):
- ret_label.append(
- [self.build_label for _ in ret_images])
- elif isinstance(self.build_label[k], list):
- assert len(
- self.build_label[k]) == len(ret_images)
- ret_label.append(self.build_label[k])
- else:
- ret_label.append([
- self.build_label[k] for _ in ret_images
- ])
- else:
- if isinstance(self.build_label, str):
- ret_label.append([self.build_label])
- else:
- ret_label.append([self.build_label[k]])
- else:
- ret_data.append(self.save_npy(prefix_path, v))
- if self.is_image and self.build_html:
- html_prefix = prefix + '_probe.html'
- html_file = self.save_html(html_prefix, ret_data,
- ret_label)
- return {'ori_file': ret_data, 'html': html_file}
- else:
- return {'ori_file': ret_data}
- else:
- return self.data
- else:
+ ret_one_data, ret_one_label = self.save_one_media(k, v,
+ prefix_path,
+ image_postfix,
+ video_postfix,
+ rank = rank)
+ ret_data.append(ret_one_data)
+ ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
+ self.media_handler.sync()
+ self.media_handler.clear()
+ self.data = {"ret_data": ret_data, "ret_label": ret_label}
+ self.is_presave = True
+
+ def to_log(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
+ if not self.is_presave:
+ self.presave(prefix, image_postfix, video_postfix, rank = rank)
+ if not self.is_presave:
return self.data
+ if isinstance(self.data, str):
+ return self.data
+ elif isinstance(self.data, dict):
+ ret_data, ret_label = self.data["ret_data"], self.data["ret_label"]
+ if self.build_html:
+ html_prefix = prefix + '_probe.html'
+ html_file = self.save_html(html_prefix, ret_data,
+ ret_label)
+ return {'ori_file': ret_data, 'html': html_file}
+ else:
+ return {'ori_file': ret_data}
+ elif isinstance(self.data, list):
+ ret_data, ret_label = [], []
+ for one_data in self.data:
+ if isinstance(one_data, dict):
+ one_ret_data, one_ret_label = one_data["ret_data"], one_data["ret_label"]
+ ret_data.extend(one_ret_data)
+ ret_label.extend(one_ret_label)
+ elif isinstance(one_data, str):
+ ret_data.append(one_data)
+ if (self.is_image or self.is_video) and self.build_html:
+ html_prefix = prefix + '_probe.html'
+ html_file = self.save_html(html_prefix, ret_data,
+ ret_label)
+ return {'ori_file': ret_data, 'html': html_file}
+ else:
+ return {'ori_file': ret_data}
diff --git a/scepter/modules/utils/visualization.py b/scepter/modules/utils/visualization.py
new file mode 100644
index 0000000..84c5e09
--- /dev/null
+++ b/scepter/modules/utils/visualization.py
@@ -0,0 +1,257 @@
+# -*- coding: utf-8 -*-
+from enum import Enum
+
+
+class Media(Enum):
+ TEXT = 1
+ IMAGE = 2
+ VIDEO = 3
+ AUDIO = 4
+
+
+class HtmlVisualization(object):
+ def __init__(
+ self,
+ allow_annotation=False,
+ slice_size=1000,
+ align='center',
+ width_scale='60%',
+ title='Visualization',
+ ):
+ self.content_list = []
+ self.rows_meta = []
+ self.allow_annotation = allow_annotation
+ self.slice_size = slice_size
+ self.align = align
+ self.width_scale = width_scale
+ self.title = title
+ self.html_start = ''
+ self.html_head = f'{title}'
+ self.html_style = '''
+
+
+ '''.replace('{width_scale}',
+ self.width_scale).replace('{align}', self.align)
+ self.html_body = '{BODY}\n'
+ self.html_end = ''
+ self.html_script = '''
+
+ '''
+ self.label_button = (
+ '| ' +
+ ""
+ + ' |
')
+
+ def format_col(self,
+ content='',
+ label='',
+ type=Media.TEXT,
+ content_height=400,
+ content_width=600):
+ if type == Media.TEXT:
+ ret_str = ' | \n'
+ sec_ret_str = f'{label} | \n'
+ return [ret_str, sec_ret_str]
+ elif type == Media.IMAGE:
+ ret_str = f' | \n'
+ sec_ret_str = f'{label} | \n'
+ return [ret_str, sec_ret_str]
+ elif type == Media.VIDEO:
+ ret_str = ' | \n'
+ sec_ret_str = f'{label} | \n'
+ return [ret_str, sec_ret_str]
+ elif type == Media.AUDIO:
+ ret_str = f' | \n'
+ sec_ret_str = f'{label} | \n'
+ return [ret_str, sec_ret_str]
+ else:
+ raise NotImplementedError
+
+ def format_row(self):
+ sample_id = 0
+ all_sample_html = []
+ for one_content, one_row_meta in zip(self.content_list,
+ self.rows_meta):
+ one_row_str = ''
+ if self.allow_annotation:
+ one_row_str = f''
+ all_sample_html.append(one_row_str)
+ sample_id += 1
+
+ return '\n'.join(all_sample_html)
+
+ def add_record(self,
+ content='',
+ label='',
+ type=Media.TEXT,
+ row_id=1,
+ col_id=1,
+ annotation_meta=None,
+ content_height=None,
+ content_width=None):
+ if row_id >= len(self.content_list):
+ self.content_list.append([])
+ self.rows_meta.append([])
+ if row_id != len(self.content_list) - 1:
+ raise RuntimeError(
+ 'row_id should be next number of the last row_id.')
+ if col_id > len(self.content_list[row_id]):
+ raise RuntimeError(
+ 'col_id should be next number of the last col_id.')
+ format_col = self.format_col(content, f"{row_id}-{col_id}: {label}",
+ type, content_height, content_width)
+
+ annotation_meta = annotation_meta if annotation_meta else ''
+ if col_id == len(self.content_list[row_id]):
+ self.content_list[row_id].append(format_col)
+ self.rows_meta[row_id].append(annotation_meta)
+ else:
+ self.content_list[row_id][col_id] = format_col
+ self.rows_meta[row_id][col_id] = annotation_meta
+
+ def save_html(self, path):
+ html_body = self.format_row()
+ ret_html_list = [
+ self.html_start, self.html_head, self.html_style,
+ self.html_body.replace('{BODY}', html_body)
+ ]
+ if self.allow_annotation:
+ ret_html_list.append(self.label_button)
+ ret_html_list.append(self.html_script)
+ ret_html_list.append(self.html_end)
+ ret_html = '\n'.join(ret_html_list)
+ with open(path, 'w') as f:
+ f.write(ret_html)
+
+
+if __name__ == '__main__':
+ from scepter.modules.utils.config import Config
+ from scepter.modules.utils.file_system import FS
+ FS.init_fs_client(Config(cfg_dict={}, load=False))
+
+ image_content_oss = '0_probe_0_[1024_2048_3].jpg'
+ content_oss = '6UTWGRG1lx08iRBx5REA01041200dzcb0E010.mp4'
+ caption = 'a little girl says hello.'
+
+ html_ins = HtmlVisualization(allow_annotation=True,
+ slice_size=1000,
+ title='Visualization',
+ width_scale='100%')
+
+ for i in range(4):
+ content_url = FS.get_url(content_oss, skip_check=True)
+ html_ins.add_record(content=content_url,
+ label='caption',
+ type=Media.VIDEO,
+ row_id=i,
+ col_id=0,
+ annotation_meta=None,
+ content_height=600,
+ content_width=None)
+ html_ins.add_record(content=caption,
+ label='caption',
+ type=Media.TEXT,
+ row_id=i,
+ col_id=1,
+ annotation_meta=None,
+ content_height=600,
+ content_width=750)
+ image_content_url = FS.get_url(image_content_oss, skip_check=True)
+ html_ins.add_record(content=image_content_url,
+ label='caption',
+ type=Media.IMAGE,
+ row_id=i,
+ col_id=2,
+ annotation_meta=None,
+ content_height=600,
+ content_width=None)
+
+ with FS.put_to('visualize.html') as local_path:
+ html_ins.save_html(local_path)
diff --git a/scepter/studio/inference/inference_manager/infer_runer.py b/scepter/studio/inference/inference_manager/infer_runer.py
index 769a043..288ae9c 100644
--- a/scepter/studio/inference/inference_manager/infer_runer.py
+++ b/scepter/studio/inference/inference_manager/infer_runer.py
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.inference.diffusion_inference import DiffusionInference
+from scepter.modules.inference.flux_inference import FluxInference
from scepter.modules.inference.largen_inference import LargenInference
from scepter.modules.inference.pixart_inference import PixArtInference
from scepter.modules.inference.sd3_inference import SD3Inference
@@ -108,6 +109,8 @@ class PipelineManager():
PipelineBuilder = PixArtInference
elif pipeline_name.startswith('SD3'):
PipelineBuilder = SD3Inference
+ elif pipeline_name.startswith('FLUX'):
+ PipelineBuilder = FluxInference
else:
PipelineBuilder = DiffusionInference
new_inference = PipelineBuilder(logger=self.logger)
diff --git a/scepter/studio/inference/inference_ui/component_names.py b/scepter/studio/inference/inference_ui/component_names.py
index 7812c1a..c119167 100644
--- a/scepter/studio/inference/inference_ui/component_names.py
+++ b/scepter/studio/inference/inference_ui/component_names.py
@@ -6,7 +6,7 @@ from scepter.modules.utils.file_system import FS
def download_image(image):
- # return None
+ # return " "
if image is not None:
client = FS.get_fs_client(image)
if client.tmp_dir.startswith('/home'):
diff --git a/scepter/studio/inference/inference_ui/diffusion_ui.py b/scepter/studio/inference/inference_ui/diffusion_ui.py
index c7146e9..691239d 100644
--- a/scepter/studio/inference/inference_ui/diffusion_ui.py
+++ b/scepter/studio/inference/inference_ui/diffusion_ui.py
@@ -4,6 +4,8 @@ import copy
import random
import gradio as gr
+
+from scepter.modules.utils.config import Config
from scepter.studio.inference.inference_ui.component_names import \
DiffusionUIName
from scepter.studio.utils.uibase import UIBase
@@ -53,11 +55,18 @@ class DiffusionUI(UIBase):
diffusion_paras = copy.deepcopy(ori_diffusion_paras)
for key in diffusion_paras:
if key.lower() in cur_default:
- diffusion_paras.get(key).DEFAULT = cur_default.get(key.lower())
- value = diffusion_paras.get(key).get('VALUES')
- if value is not None and cur_default.get(
- key.lower()) not in value:
- value.VALUES.append(cur_default.get(key.lower()))
+ cur_pa = cur_default.get(key.lower())
+ if isinstance(cur_pa, (dict, Config)):
+ diffusion_paras.get(key).DEFAULT = cur_pa.get("DEFAULT")
+ diffusion_paras.get(key).VALUES = cur_pa.get("VALUES", [])
+ diffusion_paras.get(key).VISIBLE = cur_pa.get("VISIBLE", True)
+ else:
+ diffusion_paras.get(key).DEFAULT = cur_pa
+ value = diffusion_paras.get(key).get('VALUES')
+ if value is not None and cur_default.get(
+ key.lower()) not in value:
+ value.append(cur_default.get(key.lower()))
+
return diffusion_paras
def load_all_paras(self):
@@ -75,18 +84,21 @@ class DiffusionUI(UIBase):
placeholder=self.component_names.negative_prompt_placeholder,
info=self.component_names.negative_prompt_description,
value=self.cur_paras.NEGATIVE_PROMPT.get('DEFAULT', ''),
+ visible=self.cur_paras.NEGATIVE_PROMPT.get('VISIBLE', True),
lines=2)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.prompt_prefix = gr.Textbox(
label=self.component_names.prompt_prefix,
value=self.cur_paras.PROMPT_PREFIX.get('DEFAULT', ''),
+ visible=self.cur_paras.PROMPT_PREFIX.get('VISIBLE', True),
interactive=True)
with gr.Column(scale=2):
self.sampler = gr.Dropdown(
label=self.component_names.sample,
choices=self.cur_paras.SAMPLE.get('VALUES', []),
value=self.cur_paras.SAMPLE.get('DEFAULT', ''),
+ visible=self.cur_paras.SAMPLE.get('VISIBLE', True),
interactive=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
@@ -94,6 +106,7 @@ class DiffusionUI(UIBase):
label=self.component_names.discretization,
choices=self.cur_paras.DISCRETIZATION.get('VALUES', []),
value=self.cur_paras.DISCRETIZATION.get('DEFAULT', ''),
+ visible=self.cur_paras.DISCRETIZATION.get('VISIBLE', True),
interactive=True)
self.cur_h_level_dict, default_res = self.merge_resolutions(
self.h_level_dict, self.default_resolutions)
@@ -116,6 +129,7 @@ class DiffusionUI(UIBase):
maximum=self.cur_paras.SAMPLES.get('MAX', 4),
step=1,
value=self.cur_paras.SAMPLES.get('DEFAULT', 1),
+ visible=self.cur_paras.SAMPLES.get('VISIBLE', True),
interactive=True)
with gr.Row(equal_height=True):
self.sample_steps = gr.Slider(
@@ -132,6 +146,7 @@ class DiffusionUI(UIBase):
maximum=self.cur_paras.GUIDE_SCALE.get('MAX', 10),
step=0.5,
value=self.cur_paras.GUIDE_SCALE.get('DEFAULT', 7.5),
+ visible=self.cur_paras.GUIDE_SCALE.get('VISIBLE', True),
interactive=True)
self.guide_rescale = gr.Slider(
label=self.component_names.guide_rescale,
@@ -139,6 +154,7 @@ class DiffusionUI(UIBase):
maximum=self.cur_paras.GUIDE_RESCALE.get('MAX', 1.0),
step=0.1,
value=self.cur_paras.GUIDE_RESCALE.get('DEFAULT', 0.5),
+ visible=self.cur_paras.GUIDE_RESCALE.get('VISIBLE', True),
interactive=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
diff --git a/scepter/studio/inference/inference_ui/model_manage_ui.py b/scepter/studio/inference/inference_ui/model_manage_ui.py
index adbfe9b..b1e8319 100644
--- a/scepter/studio/inference/inference_ui/model_manage_ui.py
+++ b/scepter/studio/inference/inference_ui/model_manage_ui.py
@@ -205,19 +205,26 @@ class ModelManageUI(UIBase):
value=[]),
gr.Textbox(elem_classes=cur_paras.NEGATIVE_PROMPT.get(
'VALUES', []),
- value=cur_paras.NEGATIVE_PROMPT.get('DEFAULT', '')),
+ value=cur_paras.NEGATIVE_PROMPT.get('DEFAULT', ''),
+ visible=cur_paras.NEGATIVE_PROMPT.get('VISIBLE', True)),
gr.Textbox(elem_classes=cur_paras.PROMPT_PREFIX.get(
'VALUES', []),
- value=cur_paras.PROMPT_PREFIX.get('DEFAULT', '')),
+ value=cur_paras.PROMPT_PREFIX.get('DEFAULT', ''),
+ visible=cur_paras.PROMPT_PREFIX.get('VISIBLE', True)),
gr.Dropdown(choices=[key for key in h_level_dict.keys()],
value=default_res[0]),
gr.Dropdown(choices=cur_paras.SAMPLE.get('VALUES', []),
- value=cur_paras.SAMPLE.get('DEFAULT', '')),
+ value=cur_paras.SAMPLE.get('DEFAULT', ''),
+ visible=cur_paras.SAMPLE.get('VISIBLE', True)),
gr.Dropdown(choices=cur_paras.DISCRETIZATION.get('VALUES', []),
- value=cur_paras.DISCRETIZATION.get('DEFAULT', '')),
- gr.Slider(value=cur_paras.SAMPLE_STEPS.get('DEFAULT', 30)),
- gr.Slider(value=cur_paras.GUIDE_SCALE.get('DEFAULT', 7.5)),
- gr.Slider(value=cur_paras.GUIDE_RESCALE.get('DEFAULT', 0.5)))
+ value=cur_paras.DISCRETIZATION.get('DEFAULT', ''),
+ visible=cur_paras.DISCRETIZATION.get('VISIBLE', True)),
+ gr.Slider(value=cur_paras.SAMPLE_STEPS.get('DEFAULT', 30),
+ visible=cur_paras.SAMPLE_STEPS.get('VISIBLE', True)),
+ gr.Slider(value=cur_paras.GUIDE_SCALE.get('DEFAULT', 7.5),
+ visible=cur_paras.GUIDE_SCALE.get('VISIBLE', True)),
+ gr.Slider(value=cur_paras.GUIDE_RESCALE.get('DEFAULT', 0.5),
+ visible=cur_paras.GUIDE_RESCALE.get('VISIBLE', True)))
self.diffusion_model.change(
diffusion_model_change,
diff --git a/scepter/studio/preprocess/caption_editor_ui/component_names.py b/scepter/studio/preprocess/caption_editor_ui/component_names.py
index c655b86..651d72d 100644
--- a/scepter/studio/preprocess/caption_editor_ui/component_names.py
+++ b/scepter/studio/preprocess/caption_editor_ui/component_names.py
@@ -155,6 +155,7 @@ class DatasetGalleryUIName():
self.delete_blank_dataset = 'Blank dataset is not allowed deleting.'
self.upload_image = 'Upload Target Image'
self.upload_src_image = 'Upload Source Image'
+ self.upload_src_mask = 'Mask Image'
self.upload_image_btn = '\U00002714' # ✔️
self.cancel_upload_btn = '\U00002716' # ✖️
self.image_caption = 'Image Caption'
@@ -165,10 +166,12 @@ class DatasetGalleryUIName():
self.ori_caption = 'Original Caption'
self.dataset_images = f'Original Images,click{self.btn_modify} into editable mode.'
self.dataset_src_images = f'Source Images to be edited,click{self.btn_modify} into editable mode.'
+ self.dataset_src_mask = 'Source Image Mask'
self.edit_caption = f'Editable Caption,click{self.btn_modify} into editable mode.'
self.edit_dataset_images = f'Editable Images,click{self.btn_modify} into editable mode.'
- self.edit_dataset_src_images = 'Editable Source Images to be edited.'
+ self.edit_dataset_src_images = 'Editable Source Images to be edited'
+ self.edit_dataset_src_mask = 'Editable Source Images Mask'
self.ori_dataset = 'Original Data Height({}) * Width({}) and Image Format({})'
self.edit_dataset = 'Editable Data Height({}) * Width({}) and Image Format({})'
@@ -195,10 +198,18 @@ class DatasetGalleryUIName():
self.preprocess_choices = [
'Image Preprocess', 'Caption Preprocess'
]
+
+ self.preview_target_image = 'Preview Target Image'
+ self.preview_src_image = 'Preview Source Image'
+ self.preview_src_mask_image = 'Preview Source Image Mask'
+ self.preview_caption = 'Preview Caption'
+
self.image_processor_type = 'Image Preprocessors'
self.caption_processor_type = 'Caption Preprocessors'
- self.image_preprocess_btn = 'Run'
- self.caption_preprocess_btn = 'Run'
+ self.image_preprocess_btn = 'apply'
+ self.image_preview_btn = 'preview'
+ self.caption_preprocess_btn = 'apply'
+ self.caption_preview_btn = 'preview'
self.caption_update_mode = 'Caption Update Mode'
self.caption_update_choices = ['Append', 'Replace']
@@ -209,6 +220,7 @@ class DatasetGalleryUIName():
self.system_prompt = 'System Prompt'
self.max_new_tokens = 'Max New Tokens'
self.min_new_tokens = 'Min New Tokens'
+ self.use_local = 'Regional Caption'
self.num_beams = 'Beams Num'
self.repetition_penalty = 'Repetition Penalty'
self.temperature = 'Temperature'
@@ -223,6 +235,7 @@ class DatasetGalleryUIName():
self.delete_blank_dataset = '空白数据集不允许删除。'
self.upload_image = '上传目标图片'
self.upload_src_image = '上传待编辑图片'
+ self.upload_src_mask = '蒙版区域'
self.upload_image_btn = '\U00002714' # ✔️
self.cancel_upload_btn = '\U00002716' # ✖️
self.image_caption = '图片描述'
@@ -230,9 +243,11 @@ class DatasetGalleryUIName():
self.btn_modify = '\U0001F4DD' # 📝
self.dataset_images = f'图片数据,点击{self.btn_modify}进入编辑模式'
self.dataset_src_images = f'待编辑图片数据,点击{self.btn_modify}进入编辑模式'
+ self.dataset_src_mask = '蒙版区域'
self.edit_dataset_images = '可编辑图片数据'
self.edit_dataset_src_images = '可编辑待编辑图片数据'
+ self.edit_dataset_src_mask = '可编辑待编辑图片数据蒙版'
self.btn_delete = '\U0001f5d1' # 🗑️
self.btn_add = '\U00002795' # ➕
@@ -262,10 +277,16 @@ class DatasetGalleryUIName():
f'点击{self.btn_reset_edit}重置数据,'
f'修改编辑范围可以批量编辑不同范围的数据。')
self.preprocess_choices = ['图像预处理', '描述生成']
+ self.preview_target_image = '预览图片'
+ self.preview_src_image = '预览原图'
+ self.preview_src_mask_image = '预览蒙版'
+ self.preview_caption = '预览描述'
self.image_processor_type = '图像预处理器'
self.caption_processor_type = '描述生成器'
- self.image_preprocess_btn = '运行'
- self.caption_preprocess_btn = '运行'
+ self.image_preprocess_btn = '应用'
+ self.image_preview_btn = '预览'
+ self.caption_preprocess_btn = '应用'
+ self.caption_preview_btn = '预览'
self.caption_update_mode = '描述更新方式'
self.caption_update_choices = ['追加', '替换']
self.used_device = '使用设备'
@@ -275,6 +296,7 @@ class DatasetGalleryUIName():
self.system_prompt = '系统提示'
self.max_new_tokens = '描述最大长度'
self.min_new_tokens = '描述最小长度'
+ self.use_local = '局部描述'
self.num_beams = 'Beams数'
self.repetition_penalty = '重复惩罚'
self.temperature = '温度系数'
diff --git a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py
index fb6ba45..26674a4 100644
--- a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py
+++ b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py
@@ -25,6 +25,11 @@ def wget_file(file_url, save_file):
if 'oss' in file_url:
file_url = file_url.split('?')[0]
local_path, _ = FS.map_to_local(save_file)
+ if os.path.exists(local_path):
+ try:
+ os.remove(local_path)
+ except:
+ pass
res = os.popen(f"wget -c '{file_url.strip()}' -O '{local_path.strip()}'")
res.readlines()
FS.put_object_from_local_file(local_path, save_file)
diff --git a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py
index 65a093c..99ca7b0 100644
--- a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py
+++ b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py
@@ -2,11 +2,13 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
from __future__ import annotations
+import json
import os.path
import time
import gradio as gr
import imagehash
+# from gradio.processing_utils import encode_pil_to_base64
from PIL import Image
from scepter.modules.utils.directory import get_md5
from scepter.modules.utils.file_system import FS
@@ -21,6 +23,11 @@ from scepter.studio.utils.uibase import UIBase
class DatasetGalleryUI(UIBase):
def __init__(self, cfg, is_debug=False, language='en', create_ins=None):
+ self.work_dir = cfg.WORK_DIR
+ self.local_work_dir, _ = FS.map_to_local(self.work_dir)
+ self.cache = os.path.join(self.local_work_dir, 'mask_cache')
+ os.system(f'rm -rf {self.cache}')
+ os.makedirs(self.cache, exist_ok=True)
self.selected_path = ''
self.selected_index = -1
self.selected_index_prev = -1
@@ -94,37 +101,52 @@ class DatasetGalleryUI(UIBase):
self.default_edit_image_height,
self.default_edit_image_width,
self.default_edit_image_format))
- with gr.Row():
- self.edit_caption = gr.Textbox(
- label=self.component_names.edit_caption,
- placeholder='',
- value=self.default_edit_caption,
- lines=4,
- autoscroll=False,
- interactive=True,
- visible=False)
+
with gr.Row(equal_height=True):
with gr.Column(scale=1, min_width=0,
visible=False) as self.edit_src_panel:
- self.edit_src_gl_dataset_images = gr.Gallery(
- label=self.component_names.edit_dataset_src_images,
- elem_id='dataset_tag_editor_dataset_src_gallery',
- value=self.default_image_list,
- selected_index=self.default_select_index,
- columns=4,
- visible=False,
- interactive=False)
+ with gr.Row():
+ self.edit_src_gl_dataset_images = gr.Gallery(
+ label=self.component_names.
+ edit_dataset_src_images,
+ elem_id=
+ 'dataset_tag_editor_dataset_src_gallery',
+ value=self.default_image_list,
+ selected_index=self.default_select_index,
+ preview=True,
+ allow_preview=True,
+ columns=4,
+ visible=False,
+ interactive=False)
+ with gr.Row():
+ self.edit_src_mask = gr.Image(
+ label=self.component_names.
+ edit_dataset_src_mask,
+ height=240,
+ sources=[],
+ interactive=True,
+ type='pil')
with gr.Column(scale=1, min_width=0):
- self.edit_gl_dataset_images = gr.Gallery(
- label=self.component_names.edit_dataset_images,
- elem_id='dataset_tag_editor_dataset_gallery',
- value=self.default_image_list,
- selected_index=self.default_select_index,
- object_fit='fill',
- preview=True,
- columns=4,
- visible=False,
- interactive=False)
+ with gr.Row():
+ self.edit_gl_dataset_images = gr.Gallery(
+ label=self.component_names.edit_dataset_images,
+ elem_id='dataset_tag_editor_dataset_gallery',
+ value=self.default_image_list,
+ selected_index=self.default_select_index,
+ columns=4,
+ preview=True,
+ allow_preview=True,
+ visible=False,
+ interactive=False)
+ with gr.Row():
+ self.edit_caption = gr.Textbox(
+ label=self.component_names.edit_caption,
+ placeholder='',
+ value=self.default_edit_caption,
+ lines=8,
+ autoscroll=False,
+ interactive=True,
+ visible=False)
# with gr.Column(ariant='panel', scale=2, min_width=0):
with gr.Column(variant='panel', min_width=0):
@@ -134,25 +156,29 @@ class DatasetGalleryUI(UIBase):
self.default_image_height,
self.default_image_width,
self.default_image_format))
- with gr.Row():
- self.ori_caption = gr.Textbox(
- label=self.component_names.ori_caption,
- placeholder='',
- value=self.default_ori_caption,
- lines=4,
- autoscroll=False,
- interactive=False)
+
with gr.Row(equal_height=True):
with gr.Column(scale=1, min_width=0,
visible=False) as self.src_panel:
- self.src_gl_dataset_images = gr.Gallery(
- label=self.component_names.dataset_src_images,
- elem_id='dataset_tag_editor_dataset_src_gallery',
- value=self.default_image_list,
- selected_index=self.default_select_index,
- columns=4,
- visible=False,
- interactive=False)
+ with gr.Row():
+ self.src_gl_dataset_images = gr.Gallery(
+ label=self.component_names.dataset_src_images,
+ elem_id=
+ 'dataset_tag_editor_dataset_src_gallery',
+ preview=True,
+ allow_preview=True,
+ value=self.default_image_list,
+ selected_index=self.default_select_index,
+ columns=4,
+ visible=False,
+ interactive=False)
+ with gr.Row():
+ self.src_mask = gr.Image(
+ label=self.component_names.dataset_src_mask,
+ height=240,
+ sources=[],
+ interactive=True,
+ type='pil')
with gr.Column(scale=1, min_width=0):
with gr.Row():
self.gl_dataset_images = gr.Gallery(
@@ -164,11 +190,21 @@ class DatasetGalleryUI(UIBase):
preview=True,
allow_preview=True,
object_fit='fill')
+ with gr.Row():
+ self.ori_caption = gr.Textbox(
+ label=self.component_names.ori_caption,
+ placeholder='',
+ value=self.default_ori_caption,
+ lines=8,
+ autoscroll=False,
+ interactive=False)
+ # with gr.Column(variant='panel', scale=2, min_width=0):
+ # with gr.Row(visible=False) as :
with gr.Row(visible=False) as self.edit_setting_panel:
self.sys_log = gr.Markdown(
self.component_names.system_log.format(''))
- with (gr.Row()):
+ with gr.Row():
with gr.Column(variant='panel',
visible=False,
scale=1,
@@ -208,7 +244,7 @@ class DatasetGalleryUI(UIBase):
value=None)
with gr.Column(variant='panel',
visible=False,
- scale=2,
+ scale=1,
min_width=0) as self.upload_panel:
with gr.Row():
self.upload_src_image_info = gr.Markdown(value='',
@@ -217,26 +253,36 @@ class DatasetGalleryUI(UIBase):
self.upload_image_info = gr.Markdown(value='')
with gr.Row():
- self.upload_caption = gr.Textbox(
- label=self.component_names.image_caption,
- placeholder='',
- value='',
- lines=5)
- with gr.Row():
- with gr.Column(
+ with gr.Row(
variant='panel',
- scale=1,
- min_width=0,
visible=self.default_dataset_type ==
'scepter_img2img') as self.upload_src_image_panel:
- self.upload_src_image = gr.Image(
- label=self.component_names.upload_src_image,
- type='pil')
- with gr.Column(variant='panel', scale=1, min_width=0):
+ with gr.Column(scale=1):
+ self.upload_src_image = gr.Image(
+ label=self.component_names.upload_src_image,
+ sources=['upload'],
+ interactive=True,
+ type='pil')
+ with gr.Column(scale=1):
+ self.upload_src_mask = gr.Image(
+ label=self.component_names.upload_src_mask,
+ sources=['upload'],
+ type='pil',
+ height=240,
+ interactive=False)
+ with gr.Row():
+ with gr.Column(scale=1):
self.upload_image = gr.Image(
label=self.component_names.upload_image,
sources=['upload'],
type='pil')
+ with gr.Column(scale=1):
+ self.upload_caption = gr.Textbox(
+ label=self.component_names.image_caption,
+ autoscroll=True,
+ placeholder='',
+ value='',
+ lines=8)
with gr.Row():
with gr.Column(min_width=0):
self.upload_button = gr.Button(
@@ -244,10 +290,52 @@ class DatasetGalleryUI(UIBase):
with gr.Column(min_width=0):
self.cancel_button = gr.Button(
value=self.component_names.cancel_upload_btn)
+ with gr.Row():
+ self.upload_preprocess_checkbox = gr.CheckboxGroup(
+ show_label=False,
+ choices=self.component_names.preprocess_choices,
+ value=None)
with gr.Column(variant='panel',
visible=False,
- scale=2,
+ scale=3,
min_width=0) as self.preprocess_panel:
+ with gr.Row(variant='panel') as self.preview_panel:
+ with gr.Column(scale=1,
+ visible=False) as self.preview_src_panel:
+ self.preview_src_image = gr.ImageMask(
+ label=self.component_names.preview_src_image,
+ sources=[],
+ layers=False,
+ type='pil',
+ interactive=True)
+ self.preview_src_image_tool = gr.State(value='sketch')
+ with gr.Column(
+ scale=1,
+ visible=False) as self.preview_src_mask_panel:
+ self.preview_src_mask_image = gr.ImageMask(
+ label=self.component_names.preview_src_mask_image,
+ sources=[],
+ layers=False,
+ type='pil',
+ interactive=True)
+ self.preview_src_mask_image_tool = gr.State(
+ value='sketch')
+ with gr.Column(scale=1):
+ self.preview_taget_image = gr.ImageMask(
+ label=self.component_names.preview_target_image,
+ sources=[],
+ layers=False,
+ type='pil',
+ interactive=True)
+ self.preview_taget_image_tool = gr.State(
+ value='sketch')
+ with gr.Column(scale=1):
+ self.preview_caption = gr.Textbox(
+ label=self.component_names.preview_caption,
+ placeholder='',
+ lines=4,
+ autoscroll=True,
+ interactive=True)
with gr.Row(
variant='panel',
visible=False,
@@ -266,6 +354,14 @@ class DatasetGalleryUI(UIBase):
image_processor_ins = self.processors_manager.get_processor(
'image',
self.processors_manager.get_default('image'))
+ with gr.Row():
+ use_mask_visible = image_processor_ins.system_para.get(
+ 'USE_MASK_VISIBLE', False)
+ self.use_mask = gr.Checkbox(
+ value=True,
+ visible=use_mask_visible,
+ interactive=True)
+
with gr.Row():
default_height_ratio = image_processor_ins.system_para.get(
'HEIGHT_RATIO', {})
@@ -290,9 +386,19 @@ class DatasetGalleryUI(UIBase):
visible=len(default_width_ratio) > 0,
interactive=True)
with gr.Row():
- self.image_preprocess_btn = gr.Button(
- value=self.component_names.
- image_preprocess_btn)
+ with gr.Column(
+ ) as self.image_preview_btn_panel:
+ self.image_preview_btn = gr.Button(
+ value=self.component_names.
+ image_preview_btn)
+ with gr.Column():
+ self.image_preprocess_btn = gr.Button(
+ value=self.component_names.
+ image_preprocess_btn)
+ with gr.Column(scale=1, min_width=0):
+ self.image_preprocess_reset_btn = gr.Button(
+ value=self.component_names.
+ btn_reset_edit)
with gr.Row(variant='panel',
visible=False) as self.caption_preprocess_panel:
with gr.Group():
@@ -342,6 +448,15 @@ class DatasetGalleryUI(UIBase):
if len(self.component_names.
caption_update_choices) > 0 else
None)
+ with gr.Row():
+ default_use_local = default_processor_ins.get_para_by_language(
+ default_processor_ins.get_language_default
+ ).get('USE_LOCAL', False)
+ self.use_local = gr.Checkbox(
+ label=self.component_names.use_local,
+ value=default_use_local,
+ visible=True,
+ interactive=True)
with gr.Accordion(
label=self.component_names.advance_setting,
open=False):
@@ -361,6 +476,7 @@ class DatasetGalleryUI(UIBase):
get_para_by_language(
default_processor_ins.
get_language_default))
+
with gr.Row():
default_max_new_tokens = default_processor_ins.get_para_by_language(
default_processor_ins.
@@ -380,6 +496,7 @@ class DatasetGalleryUI(UIBase):
visible=len(default_max_new_tokens) >
0,
interactive=True)
+
with gr.Row():
default_min_new_tokens = default_processor_ins.get_para_by_language(
default_processor_ins.
@@ -458,7 +575,10 @@ class DatasetGalleryUI(UIBase):
interactive=True)
with gr.Row():
-
+ with gr.Column(scale=1, min_width=0):
+ self.caption_preview_btn = gr.Button(
+ value=self.component_names.
+ caption_preview_btn)
with gr.Column(scale=1, min_width=0):
self.caption_preprocess_btn = gr.Button(
value=self.component_names.
@@ -497,7 +617,8 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return gr.Gallery(), gr.Gallery(), gr.Column(), gr.Text()
+ return gr.Gallery(), gr.Gallery(), gr.Image(), gr.Column(
+ ), gr.Text()
cursor = dataset_ins.cursor
image_list = [
os.path.join(dataset_ins.meta['local_work_dir'],
@@ -509,6 +630,14 @@ class DatasetGalleryUI(UIBase):
v['src_relative_path'])
for v in dataset_ins.data
]
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['src_mask_relative_path'])
+ else:
+ mask_path = None
+ ret_mask = gr.Image(value=mask_path, visible=True)
if cursor >= 0:
ret_src_gallery = gr.Gallery(label=dataset_name,
value=src_image_list,
@@ -523,24 +652,25 @@ class DatasetGalleryUI(UIBase):
else:
ret_src_gallery = gr.Gallery(visible=False)
ret_src_panel = gr.Column(visible=False)
+ ret_mask = gr.Image(visible=False)
if cursor >= 0:
return (gr.Gallery(label=dataset_name,
value=image_list,
selected_index=cursor), ret_src_gallery,
- ret_src_panel, gr.Text(value='view'))
+ ret_mask, ret_src_panel, gr.Text(value='view'))
else:
return (gr.Gallery(label=dataset_name,
value=image_list,
selected_index=None), ret_src_gallery,
- ret_src_panel, gr.Text(value='view'))
+ ret_mask, ret_src_panel, gr.Text(value='view'))
self.gallery_state.change(
change_gallery,
inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
outputs=[
self.gl_dataset_images, self.src_gl_dataset_images,
- self.src_panel, self.mode_state
+ self.src_mask, self.src_panel, self.mode_state
],
queue=False)
@@ -550,7 +680,8 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return gr.Gallery(), gr.Gallery(), gr.Textbox(), '', gr.Text()
+ return gr.Gallery(), gr.Gallery(), gr.Image(), gr.Textbox(
+ ), '', gr.Text()
dataset_ins.set_cursor(evt.index)
current_info = dataset_ins.current_record
all_number = len(dataset_ins)
@@ -570,8 +701,17 @@ class DatasetGalleryUI(UIBase):
]
ret_src_image_gl = gr.Gallery(value=src_image_list, selected_index=dataset_ins.cursor, visible=True) \
if dataset_ins.cursor >= 0 else gr.Gallery(value=[], selected_index=None)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['src_mask_relative_path'])
+ else:
+ mask_path = None
+ ret_mask = gr.Image(value=mask_path, visible=True)
else:
ret_src_image_gl = gr.Gallery(visible=False)
+ ret_mask = gr.Image(visible=False)
ret_caption = gr.Textbox(value=current_info.get('caption', ''))
ret_info = gr.Text(value=f'{dataset_ins.cursor+1}/{all_number}')
@@ -584,15 +724,15 @@ class DatasetGalleryUI(UIBase):
image_info = self.component_names.ori_dataset.format(
ret_image_height, ret_image_width, ret_image_format)
- return (ret_image_gl, ret_src_image_gl, ret_caption, image_info,
- ret_info)
+ return (ret_image_gl, ret_src_image_gl, ret_mask, ret_caption,
+ image_info, ret_info)
self.gl_dataset_images.select(
select_image,
inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
outputs=[
self.gl_dataset_images, self.src_gl_dataset_images,
- self.ori_caption, self.image_info, self.info
+ self.src_mask, self.ori_caption, self.image_info, self.info
],
queue=False)
@@ -603,8 +743,8 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return gr.Gallery(), gr.Gallery(), gr.Gallery(), gr.Gallery(
- ), gr.Textbox(), '', gr.Text()
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Gallery(),
+ gr.Gallery(), gr.Image(), gr.Textbox(), '', gr.Text())
cursor = dataset_ins.cursor_from_edit_index(evt.index)
dataset_ins.set_cursor(cursor)
@@ -624,9 +764,25 @@ class DatasetGalleryUI(UIBase):
ret_src_edit_image_gl = gr.Gallery(selected_index=dataset_ins.edit_cursor, visible=True) \
if dataset_ins.cursor >= 0 \
else gr.Gallery(value=[], selected_index=None, visible=True)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['src_mask_relative_path'])
+ edit_mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info.get(
+ 'edit_src_mask_relative_path',
+ current_info['src_mask_relative_path']))
+ else:
+ mask_path, edit_mask_path = None, None
+ ret_mask = gr.Image(value=mask_path, visible=True)
+ ret_edit_mask = gr.Image(value=edit_mask_path, visible=True)
else:
ret_src_image_gl = gr.Gallery(visible=False)
ret_src_edit_image_gl = gr.Gallery(visible=False)
+ ret_mask = gr.Image(visible=False)
+ ret_edit_mask = gr.Image(visible=False)
ret_edit_caption = gr.Textbox(value=current_info.get(
'edit_caption', ''),
@@ -645,9 +801,10 @@ class DatasetGalleryUI(UIBase):
ret_edit_image_format)
ret_info = gr.Text(value=f'{dataset_ins.cursor+1}/{all_number}')
- return (ret_image_gl, ret_src_image_gl, ret_edit_image_gl,
- ret_src_edit_image_gl, ret_edit_caption, edit_image_info,
- ret_info)
+
+ return (ret_image_gl, ret_src_image_gl, ret_mask,
+ ret_edit_image_gl, ret_src_edit_image_gl, ret_edit_mask,
+ ret_edit_caption, edit_image_info, ret_info)
self.edit_gl_dataset_images.select(
edit_select_image,
@@ -657,7 +814,8 @@ class DatasetGalleryUI(UIBase):
],
outputs=[
self.gl_dataset_images, self.src_gl_dataset_images,
- self.edit_gl_dataset_images, self.edit_src_gl_dataset_images,
+ self.src_mask, self.edit_gl_dataset_images,
+ self.edit_src_gl_dataset_images, self.edit_src_mask,
self.edit_caption, self.edit_image_info, self.info
],
queue=False)
@@ -669,15 +827,23 @@ class DatasetGalleryUI(UIBase):
if dataset_ins is None:
return gr.Gallery(), gr.Gallery(), gr.Textbox(), '', gr.Text()
dataset_ins.set_cursor(evt.index)
+
ret_image_gl = gr.Gallery(selected_index=dataset_ins.cursor) if dataset_ins.cursor >= 0 \
else gr.Gallery(value=[], selected_index=None)
-
- return ret_image_gl
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['src_mask_relative_path'])
+ else:
+ mask_path = None
+ ret_mask = gr.Image(value=mask_path)
+ return ret_image_gl, ret_mask
self.src_gl_dataset_images.select(
src_select_image,
inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
- outputs=[self.gl_dataset_images],
+ outputs=[self.gl_dataset_images, self.src_mask],
queue=False)
def edit_src_select_image(dataset_type, dataset_name, mode_state,
@@ -686,13 +852,22 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return gr.Gallery()
+ return gr.Gallery(), gr.Image()
cursor = dataset_ins.cursor_from_edit_index(evt.index)
dataset_ins.set_cursor(cursor)
dataset_ins.edit_cursor = evt.index
ret_edit_image_gl = gr.Gallery(
selected_index=dataset_ins.edit_cursor)
- return ret_edit_image_gl
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ edit_mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info.get('edit_src_mask_relative_path',
+ current_info['src_mask_relative_path']))
+ else:
+ edit_mask_path = None
+ ret_edit_mask = gr.Image(value=edit_mask_path)
+ return ret_edit_image_gl, ret_edit_mask
self.edit_src_gl_dataset_images.select(
edit_src_select_image,
@@ -700,7 +875,7 @@ class DatasetGalleryUI(UIBase):
create_dataset.dataset_type, create_dataset.dataset_name,
self.mode_state
],
- outputs=[self.edit_gl_dataset_images],
+ outputs=[self.edit_gl_dataset_images, self.edit_src_mask],
queue=False)
# target image gallery changed
@@ -769,7 +944,7 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return gr.Gallery(), gr.Text(), ''
+ return gr.Gallery(), gr.Text(), '', gr.Gallery(), gr.Image()
for idx, one_data in enumerate(dataset_ins.edit_samples):
current_cursor = dataset_ins.cursor_from_edit_index(idx)
dataset_ins.edit_samples[idx][
@@ -799,6 +974,19 @@ class DatasetGalleryUI(UIBase):
'edit_src_height'] = dataset_ins.data[current_cursor][
'src_height']
+ dataset_ins.edit_samples[idx][
+ 'edit_src_mask_relative_path'] = dataset_ins.data[
+ current_cursor]['src_mask_relative_path']
+ dataset_ins.edit_samples[idx][
+ 'edit_src_mask_path'] = dataset_ins.data[
+ current_cursor]['src_mask_path']
+ dataset_ins.edit_samples[idx][
+ 'edit_src_mask_width'] = dataset_ins.data[
+ current_cursor]['src_mask_width']
+ dataset_ins.edit_samples[idx][
+ 'edit_src_mask_height'] = dataset_ins.data[
+ current_cursor]['src_mask_height']
+
image_list = [
os.path.join(dataset_ins.meta['local_work_dir'],
v.get('edit_relative_path', v['relative_path']))
@@ -818,16 +1006,42 @@ class DatasetGalleryUI(UIBase):
ret_edit_image_height, ret_edit_image_width,
ret_edit_image_format)
+ if dataset_type == 'scepter_img2img':
+ src_image_list = [
+ os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ v.get('edit_src_relative_path',
+ v['src_relative_path']))
+ for v in dataset_ins.edit_samples
+ ]
+ ret_src_edit_image_gl = gr.Gallery(value=src_image_list,
+ selected_index=dataset_ins.edit_cursor, visible=True) \
+ if dataset_ins.cursor >= 0 \
+ else gr.Gallery(value=src_image_list, selected_index=None, visible=True)
+ if len(current_info) > 0:
+ edit_mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info.get(
+ 'edit_src_mask_relative_path',
+ current_info['src_mask_relative_path']))
+ else:
+ edit_mask_path = None
+ ret_edit_mask = gr.Image(value=edit_mask_path, visible=True)
+ else:
+ ret_src_edit_image_gl = gr.Gallery()
+ ret_edit_mask = gr.Image()
+
return (gr.Gallery(value=image_list, visible=True),
gr.Textbox(value=current_info.get('edit_caption', '')),
- edit_image_info)
+ edit_image_info, ret_src_edit_image_gl, ret_edit_mask)
self.btn_reset_edit.click(
reset_edit,
inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
outputs=[
self.edit_gl_dataset_images, self.edit_caption,
- self.edit_image_info
+ self.edit_image_info, self.edit_src_gl_dataset_images,
+ self.edit_src_mask
],
queue=False)
@@ -837,14 +1051,14 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return (gr.Gallery(), gr.Gallery(), gr.Textbox(), gr.Text(),
- gr.Text(), gr.Text(),
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Textbox(),
+ gr.Text(), gr.Text(), gr.Text(),
self.component_names.system_log.format('None'),
gr.Text())
is_flg, msg = dataset_ins.apply_changes()
if not is_flg:
- return (gr.Gallery(), gr.Gallery(), gr.Textbox(), gr.Text(),
- gr.Text(), gr.Text(),
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Textbox(),
+ gr.Text(), gr.Text(), gr.Text(),
self.component_names.system_log.format(msg), gr.Text())
image_list = [
os.path.join(dataset_ins.local_work_dir, v['relative_path'])
@@ -869,12 +1083,21 @@ class DatasetGalleryUI(UIBase):
value=src_image_list,
selected_index=dataset_ins.cursor,
visible=True)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['src_mask_relative_path'])
+ else:
+ mask_path = None
+ ret_mask = gr.Image(value=mask_path, visible=True)
else:
ret_src_image_gl = gr.Gallery(visible=False)
+ ret_mask = gr.Image(visible=False)
return (gr.Gallery(value=image_list,
selected_index=dataset_ins.cursor),
- ret_src_image_gl,
+ ret_src_image_gl, ret_mask,
gr.Textbox(value=current_record.get('caption', '')),
image_info, self.component_names.system_log.format(''),
gr.Text(value='view'))
@@ -884,12 +1107,15 @@ class DatasetGalleryUI(UIBase):
inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
outputs=[
self.gl_dataset_images, self.src_gl_dataset_images,
- self.ori_caption, self.image_info, self.sys_log,
+ self.src_mask, self.ori_caption, self.image_info, self.sys_log,
self.mode_state
],
queue=False)
- def preprocess_box_change(preprocess_checkbox):
+ def preprocess_box_change(preprocess_checkbox, dataset_name,
+ dataset_type, preview_src_image_tool,
+ preview_src_mask_image_tool,
+ preview_taget_image_tool):
image_proc_status, caption_proc_status = False, False
reverse_status = {
v: id
@@ -901,31 +1127,239 @@ class DatasetGalleryUI(UIBase):
image_proc_status = True
elif hit_status == 1:
caption_proc_status = True
- return (
- # gr.Column(visible=image_proc_status), gr.Column(
- # visible=caption_proc_status)
- gr.Column(visible=image_proc_status or caption_proc_status),
- gr.Row(visible=image_proc_status),
- gr.Row(visible=caption_proc_status))
+ if image_proc_status or caption_proc_status:
+ dataset_type = create_dataset.get_trans_dataset_type(
+ dataset_type)
+ dataset_ins = create_dataset.dataset_dict.get(
+ dataset_type, {}).get(dataset_name, None)
+ edit_index_list = dataset_ins.edit_list
+ if len(edit_index_list) > 0:
+ one_data = dataset_ins.data[edit_index_list[0]]
+ else:
+ one_data = {}
+ if dataset_type == 'scepter_img2img':
+ src_image_path = os.path.join(
+ dataset_ins.local_work_dir,
+ one_data['edit_src_relative_path'])
+ # if preview_src_image_tool == 'sketch':
+ # image = Image.open(src_image_path)
+ # w, h = image.size
+ # ret_src_image = gr.Image(value={
+ # 'image': encode_pil_to_base64(image),
+ # 'mask': default_mask(w, h)
+ # }, visible=True)
+ # else:
+ ret_src_image = gr.Image(value=src_image_path,
+ visible=True)
+ src_mask_path = os.path.join(
+ dataset_ins.local_work_dir,
+ one_data['edit_src_mask_relative_path'])
+ # if preview_src_mask_image_tool == 'sketch':
+ # image = Image.open(src_mask_path)
+ # w, h = image.size
+ # ret_src_mask = gr.Image(value={
+ # 'image': encode_pil_to_base64(image),
+ # 'mask': default_mask(w, h)
+ # }, visible=True)
+ # else:
+ ret_src_mask = gr.Image(value=src_mask_path, visible=True)
+ ret_src_panel = gr.Column(visible=True)
+ ret_src_mask_panel = gr.Column(visible=True)
+ else:
+ ret_src_image = gr.Image(visible=False)
+ ret_src_mask = gr.Image(visible=False)
+ ret_src_panel = gr.Column(visible=False)
+ ret_src_mask_panel = gr.Column(visible=False)
+ image_path = os.path.join(dataset_ins.local_work_dir,
+ one_data['edit_relative_path'])
+ # if preview_taget_image_tool == 'sketch':
+ # image = Image.open(image_path)
+ # w, h = image.size
+ # ret_target_image = gr.Image(value={
+ # 'image': encode_pil_to_base64(image),
+ # 'mask': default_mask(w, h)
+ # })
+ # else:
+ ret_target_image = gr.Image(value=image_path)
+ ret_caption = gr.Textbox(value=one_data['edit_caption'])
+ ret_panel = gr.Row(visible=True)
+ else:
+ ret_src_image = gr.Image()
+ ret_src_mask = gr.Image()
+ ret_target_image = gr.Image()
+ ret_caption = gr.Textbox()
+ ret_panel = gr.Row(visible=False)
+ ret_src_panel = gr.Column(visible=False)
+ ret_src_mask_panel = gr.Column(visible=False)
+ ret_image_preprocess_method = gr.Dropdown(
+ visible=image_proc_status)
+ return (gr.Column(
+ visible=image_proc_status or caption_proc_status),
+ gr.Row(visible=image_proc_status),
+ gr.Row(visible=caption_proc_status), ret_panel,
+ ret_src_image, ret_src_panel, ret_src_mask,
+ ret_src_mask_panel, ret_target_image, ret_caption,
+ ret_image_preprocess_method)
- self.preprocess_checkbox.change(preprocess_box_change,
- inputs=[self.preprocess_checkbox],
- outputs=[
- self.preprocess_panel,
- self.image_preprocess_panel,
- self.caption_preprocess_panel
- ],
- queue=False)
+ self.preprocess_checkbox.change(
+ preprocess_box_change,
+ inputs=[
+ self.preprocess_checkbox, create_dataset.dataset_name,
+ create_dataset.dataset_type, self.preview_src_image_tool,
+ self.preview_src_mask_image_tool, self.preview_taget_image_tool
+ ],
+ outputs=[
+ self.preprocess_panel,
+ self.image_preprocess_panel,
+ self.caption_preprocess_panel,
+ # gr.Row()
+ self.preview_panel,
+ # gr.Image()
+ self.preview_src_image,
+ # gr.Column()
+ self.preview_src_panel,
+ # gr.Image()
+ self.preview_src_mask_image,
+ # gr.Column()
+ self.preview_src_mask_panel,
+ # gr.Image()
+ self.preview_taget_image,
+ # gr.TextBox()
+ self.preview_caption,
+ self.image_preprocess_method
+ ],
+ queue=False)
- def preprocess_image(mode_state, preprocess_method, upload_image,
- upload_src_image, upload_caption, height_ratio,
- width_ratio, dataset_type, dataset_name):
+ self.image_preprocess_reset_btn.click(
+ preprocess_box_change,
+ inputs=[
+ self.preprocess_checkbox, create_dataset.dataset_name,
+ create_dataset.dataset_type, self.preview_src_image_tool,
+ self.preview_src_mask_image_tool, self.preview_taget_image_tool
+ ],
+ outputs=[
+ self.preprocess_panel,
+ self.image_preprocess_panel,
+ self.caption_preprocess_panel,
+ # gr.Row()
+ self.preview_panel,
+ # gr.Image()
+ self.preview_src_image,
+ # gr.Column()
+ self.preview_src_panel,
+ # gr.Image()
+ self.preview_src_mask_image,
+ # gr.Column()
+ self.preview_src_mask_panel,
+ # gr.Image()
+ self.preview_taget_image,
+ # gr.TextBox()
+ self.preview_caption,
+ self.image_preprocess_method
+ ],
+ )
+
+ def upload_preprocess_box_change(preprocess_checkbox, dataset_type,
+ upload_src_image, upload_src_mask,
+ upload_image, upload_caption):
+ dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
+ image_proc_status, caption_proc_status = False, False
+ reverse_status = {
+ v: id
+ for id, v in enumerate(self.component_names.preprocess_choices)
+ }
+ for value in preprocess_checkbox:
+ hit_status = reverse_status[value]
+ if hit_status == 0:
+ image_proc_status = True
+ elif hit_status == 1:
+ caption_proc_status = True
+ if image_proc_status or caption_proc_status:
+ if dataset_type == 'scepter_img2img':
+ ret_src_image = gr.Image(value=upload_src_image,
+ visible=True)
+ ret_src_mask = gr.Image(value=upload_src_mask,
+ sources=['upload'],
+ visible=True)
+ ret_src_panel = gr.Column(visible=True)
+ ret_src_mask_panel = gr.Column(visible=True)
+ else:
+ ret_src_image = gr.Image(visible=False)
+ ret_src_mask = gr.Image(visible=False)
+ ret_src_panel = gr.Column(visible=False)
+ ret_src_mask_panel = gr.Column(visible=False)
+
+ ret_target_image = gr.Image(value=upload_image)
+ ret_caption = gr.Textbox(value=upload_caption)
+ ret_panel = gr.Row(visible=True)
+ else:
+ ret_src_image = gr.Image()
+ ret_src_mask = gr.Image()
+ ret_target_image = gr.Image()
+ ret_caption = gr.Textbox()
+ ret_panel = gr.Row(visible=False)
+ ret_src_panel = gr.Column(visible=False)
+ ret_src_mask_panel = gr.Column(visible=False)
+ return (gr.Column(
+ visible=image_proc_status or caption_proc_status),
+ gr.Row(visible=image_proc_status),
+ gr.Row(visible=caption_proc_status), ret_panel,
+ ret_src_image, ret_src_panel, ret_src_mask,
+ ret_src_mask_panel, ret_target_image, ret_caption)
+
+ self.upload_preprocess_checkbox.change(
+ upload_preprocess_box_change,
+ inputs=[
+ self.upload_preprocess_checkbox, create_dataset.dataset_type,
+ self.upload_src_image, self.upload_src_mask, self.upload_image,
+ self.upload_caption
+ ],
+ outputs=[
+ self.preprocess_panel,
+ self.image_preprocess_panel,
+ self.caption_preprocess_panel,
+ # gr.Row()
+ self.preview_panel,
+ # gr.Image()
+ self.preview_src_image,
+ # gr.Column()
+ self.preview_src_panel,
+ # gr.Image()
+ self.preview_src_mask_image,
+ # gr.Column()
+ self.preview_src_mask_panel,
+ # gr.Image()
+ self.preview_taget_image,
+ # gr.TextBox()
+ self.preview_caption
+ ],
+ queue=False)
+
+ def preprocess_image(
+ mode_state,
+ preprocess_method,
+ upload_image,
+ upload_src_image,
+ upload_src_mask,
+ upload_caption,
+ # preview_info
+ preview_image,
+ preview_src_image,
+ preview_src_mask,
+ preview_caption,
+ # edit_info
+ use_mask,
+ height_ratio,
+ width_ratio,
+ dataset_type,
+ dataset_name):
dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return (gr.Gallery(), gr.Gallery(), '', gr.Image(), gr.Image(),
- '', '', self.component_names.system_log.format('None'))
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), '', gr.Image(),
+ gr.Image(), gr.Image(), '', '',
+ self.component_names.system_log.format('None'))
if hasattr(manager, 'inference'):
for k, v in manager.inference.pipe_manager.pipeline_level_modules.items(
):
@@ -935,38 +1369,30 @@ class DatasetGalleryUI(UIBase):
'image', preprocess_method)
if processor_ins is None:
sys_log = 'Current processor is illegal'
- return (gr.Gallery(), gr.Gallery(), '', gr.Image(),
- gr.Image(), '', '',
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), '', gr.Image(),
+ gr.Image(), gr.Image(), '', '',
self.component_names.system_log.format(sys_log))
is_flag, msg = processor_ins.load_model()
if not is_flag:
sys_log = f'Load processor failed: {msg}'
- return (gr.Gallery(), gr.Gallery(), '', gr.Image(),
- gr.Image(), '', '',
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), '', gr.Image(),
+ gr.Image(), gr.Image(), '', '',
self.component_names.system_log.format(sys_log))
+
+ if not dataset_type == 'scepter_img2img':
+ preview_src_image = None
+ preview_src_mask = None
+
if mode_state == 'edit':
save_folders = 'edit_images'
os.makedirs(os.path.join(dataset_ins.meta['local_work_dir'],
save_folders),
exist_ok=True)
- edit_index_list = dataset_ins.edit_list
- for index in edit_index_list:
- one_data = dataset_ins.data[index]
- # process target image
- relative_image_path = one_data.get(
- 'edit_relative_path', one_data['relative_path'])
- file_name, surfix = os.path.splitext(relative_image_path)
-
- input_image = os.path.join(
- dataset_ins.meta['local_work_dir'],
- relative_image_path)
- output_image = processor_ins(
- Image.open(input_image).convert('RGB'),
- height_ratio=height_ratio,
- width_ratio=width_ratio)
-
- now_time = int(time.time())
+ def save_image(saved_image, ori_relative_image_path):
+ file_name, surfix = os.path.splitext(
+ ori_relative_image_path)
+ now_time = int(time.time() * 100)
save_file_name = f'{os.path.basename(file_name)}_{now_time}'
save_relative_image_path = os.path.join(
save_folders, f'{get_md5(save_file_name)}{surfix}')
@@ -975,58 +1401,105 @@ class DatasetGalleryUI(UIBase):
local_save_image_math = os.path.join(
dataset_ins.meta['local_work_dir'],
save_relative_image_path)
- output_image.save(local_save_image_math)
+ saved_image.save(local_save_image_math)
FS.put_object_from_local_file(local_save_image_math,
save_image_path)
+ return save_image_path, save_relative_image_path
- dataset_ins.data[index][
- 'edit_relative_path'] = save_relative_image_path
- dataset_ins.data[index][
- 'edit_image_path'] = save_image_path
- dataset_ins.data[index]['edit_width'] = output_image.size[
- 0]
- dataset_ins.data[index]['edit_height'] = output_image.size[
- 1]
-
+ def proc_sample_fn(now_data):
+ # process target image
+ relative_image_path = now_data.get(
+ 'edit_relative_path', now_data['relative_path'])
+ target_image = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ relative_image_path)
+ target_image = Image.open(target_image).convert('RGB')
if dataset_type == 'scepter_img2img':
- # process target image
- src_relative_image_path = one_data.get(
+ # process src image
+ src_relative_image_path = now_data.get(
'edit_src_relative_path',
- one_data['src_relative_path'])
- file_name, surfix = os.path.splitext(
- src_relative_image_path)
+ now_data['src_relative_path'])
+ src_image = Image.open(
+ os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ src_relative_image_path)).convert('RGB')
- input_image = os.path.join(
- dataset_ins.meta['local_work_dir'],
- src_relative_image_path)
- output_image = processor_ins(
- Image.open(input_image).convert('RGB'),
- height_ratio=height_ratio,
- width_ratio=width_ratio)
+ src_mask_relative_path = now_data.get(
+ 'edit_src_mask_relative_path',
+ now_data['src_mask_relative_path'])
+ src_mask = Image.open(
+ os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ src_mask_relative_path)).convert('RGB')
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ src_image = None
+ src_relative_image_path = None
+ src_mask = None
+ src_mask_relative_path = None
+ prev_src_image = None
+ prev_src_mask = None
- now_time = int(time.time())
- save_file_name = f'{os.path.basename(file_name)}_{now_time}'
- save_src_relative_image_path = os.path.join(
- save_folders, f'{get_md5(save_file_name)}{surfix}')
- save_src_image_path = os.path.join(
- dataset_ins.meta['work_dir'],
- save_src_relative_image_path)
- local_save_src_image_math = os.path.join(
- dataset_ins.meta['local_work_dir'],
- save_src_relative_image_path)
- output_image.save(local_save_src_image_math)
- FS.put_object_from_local_file(
- local_save_src_image_math, save_src_image_path)
+ kwargs = {
+ 'width_ratio': width_ratio,
+ 'height_ratio': height_ratio,
+ 'use_mask': use_mask,
+ 'src_image': prev_src_image,
+ 'src_mask': prev_src_mask,
+ 'target_image': target_image,
+ 'caption': one_data['edit_caption'],
+ 'preview_target_image': preview_image,
+ 'preview_src_image': prev_src_image,
+ 'preview_src_mask': prev_src_mask,
+ 'preview_caption': preview_caption,
+ 'use_preview': False
+ }
+ output_data = processor_ins(**kwargs)
+ target_image = output_data['target_image'].convert('RGB')
+ src_image = output_data.get('src_image', None)
+ src_mask = output_data.get('src_mask', None)
+ ret_data = {}
+ save_image_path, save_relative_image_path = save_image(
+ target_image, relative_image_path)
+ t_w, t_h = target_image.size
+ ret_data.update({
+ 'edit_relative_path': save_relative_image_path,
+ 'edit_image_path': save_image_path,
+ 'edit_width': t_w,
+ 'edit_hight': t_h
+ })
+ if src_image is not None:
+ src_image = src_image.convert('RGB')
+ save_src_image_path, save_src_relative_path = save_image(
+ src_image, src_relative_image_path)
+ s_w, s_h = src_image.size
+ ret_data.update({
+ 'edit_src_relative_path': save_src_relative_path,
+ 'edit_src_image_path': save_src_image_path,
+ 'edit_src_width': s_w,
+ 'edit_src_height': s_h
+ })
- dataset_ins.data[index][
- 'edit_src_relative_path'] = save_src_relative_image_path
- dataset_ins.data[index][
- 'edit_src_image_path'] = save_src_image_path
- dataset_ins.data[index][
- 'edit_src_width'] = output_image.size[0]
- dataset_ins.data[index][
- 'edit_src_height'] = output_image.size[1]
+ if src_mask is not None:
+ src_mask = src_mask.convert('RGB')
+ save_src_mask_path, save_src_mask_relative_path = save_image(
+ src_mask, src_mask_relative_path)
+ sm_w, sm_h = src_mask.size
+ ret_data.update({
+ 'edit_src_mask_relative_path':
+ save_src_mask_relative_path,
+ 'edit_src_mask_path': save_src_mask_path,
+ 'edit_src_mask_width': sm_w,
+ 'edit_src_mask_height': sm_h
+ })
+ return ret_data
+
+ edit_index_list = dataset_ins.edit_list
+ for index in edit_index_list:
+ one_data = dataset_ins.data[index]
+ dataset_ins.data[index].update(proc_sample_fn(one_data))
dataset_ins.update_dataset()
image_list = [
os.path.join(
@@ -1050,7 +1523,6 @@ class DatasetGalleryUI(UIBase):
edit_image_info = self.component_names.edit_dataset.format(
ret_edit_image_height, ret_edit_image_width,
ret_edit_image_format)
-
else:
edit_image_info = ''
@@ -1065,39 +1537,70 @@ class DatasetGalleryUI(UIBase):
for v in dataset_ins.edit_samples
]
ret_src_image_gallery = gr.Gallery(value=src_image_list)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['edit_src_mask_relative_path'])
+ else:
+ mask_path = None
+ ret_edit_mask = gr.Image(value=mask_path, visible=True)
else:
ret_src_image_gallery = gr.Gallery()
-
+ ret_edit_mask = gr.Image()
ret_upload_image = gr.Image()
ret_src_upload_image = gr.Image()
+ ret_src_upload_mask = gr.Image()
+ ret_preprocess_checkbox = gr.CheckboxGroup(value=[])
ret_upload_image_info = ''
ret_upload_src_image_info = gr.Markdown(visible=False)
elif mode_state == 'add':
if isinstance(upload_image, dict):
- image = upload_image['image']
+ target_image = upload_image['image']
else:
- image = upload_image
- target_image = processor_ins(image.convert('RGB'),
- height_ratio=height_ratio,
- width_ratio=width_ratio)
- w, h = target_image.size
- ret_image_gallery = gr.Gallery()
- ret_src_image_gallery = gr.Gallery()
- ret_upload_image = gr.Image(target_image)
+ target_image = upload_image
- ret_upload_image_info = self.component_names.upload_image_info.format(
- h, w)
if dataset_type == 'scepter_img2img':
if isinstance(upload_src_image, dict):
src_image = upload_src_image['image']
else:
src_image = upload_src_image
- src_target_image = processor_ins(src_image.convert('RGB'),
- height_ratio=height_ratio,
- width_ratio=width_ratio)
-
- ret_src_upload_image = gr.Image(src_target_image)
- src_w, src_h = src_target_image.size
+ if isinstance(upload_src_mask, dict):
+ src_mask = upload_src_mask['image']
+ else:
+ src_mask = upload_src_mask
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ src_image = None
+ src_mask = None
+ prev_src_image = None
+ prev_src_mask = None
+ kwargs = {
+ 'width_ratio': width_ratio,
+ 'height_ratio': height_ratio,
+ 'use_mask': use_mask,
+ 'src_image': src_image,
+ 'src_mask': src_mask,
+ 'target_image': target_image,
+ 'caption': upload_caption,
+ 'preview_target_image': preview_image,
+ 'preview_src_image': prev_src_image,
+ 'preview_src_mask': prev_src_mask,
+ 'preview_caption': preview_caption,
+ 'use_preview': False
+ }
+ output_data = processor_ins(**kwargs)
+ target_image = output_data['target_image']
+ src_image = output_data.get('src_image', None)
+ src_mask = output_data.get('src_mask', None)
+ if src_mask is not None:
+ ret_src_upload_mask = gr.Image(src_mask)
+ else:
+ ret_src_upload_mask = gr.Image()
+ if src_image is not None:
+ ret_src_upload_image = gr.Image(src_image)
+ src_w, src_h = src_image.size
ret_upload_src_image_info = gr.Markdown(
value=self.component_names.upload_src_image_info.
format(src_h, src_w),
@@ -1105,46 +1608,77 @@ class DatasetGalleryUI(UIBase):
else:
ret_src_upload_image = gr.Image()
ret_upload_src_image_info = gr.Markdown(visible=False)
+ w, h = target_image.size
+ ret_image_gallery = gr.Gallery()
+ ret_src_image_gallery = gr.Gallery()
+ ret_edit_mask = gr.Image()
+ ret_upload_image = gr.Image(target_image)
+ ret_preprocess_checkbox = gr.CheckboxGroup(value=[])
+ ret_upload_image_info = self.component_names.upload_image_info.format(
+ h, w)
edit_image_info = ''
else:
ret_image_gallery = gr.Gallery()
ret_src_image_gallery = gr.Gallery()
ret_upload_image = gr.Image()
+ ret_src_upload_mask = gr.Image()
+ ret_edit_mask = gr.Image()
ret_src_upload_image = gr.Image()
ret_upload_image_info = ''
ret_upload_src_image_info = gr.Markdown()
+ ret_preprocess_checkbox = gr.CheckboxGroup(value=[])
edit_image_info = ''
is_flag, msg = processor_ins.unload_model()
if not is_flag:
sys_log = f'Unoad processor failed: {msg}'
- return (ret_image_gallery, edit_image_info, ret_upload_image,
- ret_src_upload_image, ret_upload_image_info,
+ return (ret_image_gallery, ret_src_image_gallery,
+ ret_edit_mask, edit_image_info, ret_upload_image,
+ ret_src_upload_image, ret_src_upload_mask,
+ ret_upload_image_info, ret_upload_src_image_info,
self.component_names.system_log.format(sys_log))
- return (ret_image_gallery, ret_src_image_gallery, edit_image_info,
- ret_upload_image, ret_src_upload_image,
- ret_upload_image_info, ret_upload_src_image_info,
+ return (ret_image_gallery, ret_src_image_gallery, ret_edit_mask,
+ edit_image_info, ret_upload_image, ret_src_upload_image,
+ ret_src_upload_mask, ret_upload_image_info,
+ ret_upload_src_image_info, ret_preprocess_checkbox,
self.component_names.system_log.format(''))
self.image_preprocess_btn.click(
preprocess_image,
inputs=[
self.mode_state, self.image_preprocess_method,
- self.upload_image, self.upload_src_image, self.upload_caption,
- self.height_ratio, self.width_ratio,
- create_dataset.dataset_type, create_dataset.dataset_name
+ self.upload_image, self.upload_src_image, self.upload_src_mask,
+ self.upload_caption, self.preview_taget_image,
+ self.preview_src_image, self.preview_src_mask_image,
+ self.preview_caption, self.use_mask, self.height_ratio,
+ self.width_ratio, create_dataset.dataset_type,
+ create_dataset.dataset_name
],
outputs=[
self.edit_gl_dataset_images, self.edit_src_gl_dataset_images,
- self.edit_image_info, self.upload_image, self.upload_src_image,
+ self.edit_src_mask, self.edit_image_info, self.upload_image,
+ self.upload_src_image, self.upload_src_mask,
self.upload_image_info, self.upload_src_image_info,
- self.sys_log
+ self.preprocess_checkbox, self.sys_log
],
queue=False)
- def image_preprocess_method_change(image_preprocess_method):
+ # def default_mask(w, h):
+ # mask = encode_pil_to_base64(Image.new('L', (w, h), 0))
+ # return mask
+
+ def image_preprocess_method_change(image_preprocess_method,
+ preview_src_image_tool,
+ preview_src_mask_tool,
+ preview_target_image_tool):
image_processor_ins = self.processors_manager.get_processor(
'image', image_preprocess_method)
+ if image_processor_ins is None:
+ return (gr.Slider(), gr.Slider(), gr.Image(), gr.Image(),
+ gr.Image(), gr.Image(), gr.Textbox(), gr.Column(),
+ preview_src_image_tool, preview_src_mask_tool,
+ preview_target_image_tool)
+
height_ratio = image_processor_ins.system_para.get(
'HEIGHT_RATIO', {})
ret_height_ratio = gr.Slider(minimum=height_ratio.get('MIN', 1),
@@ -1161,12 +1695,56 @@ class DatasetGalleryUI(UIBase):
value=width_ratio.get('VALUE', 1),
visible=len(width_ratio) > 0,
interactive=True)
- return (ret_height_ratio, ret_width_ratio)
+
+ use_mask_visible = image_processor_ins.system_para.get(
+ 'USE_MASK_VISIBLE', False)
+ use_mask = gr.Checkbox(value=True,
+ visible=use_mask_visible,
+ interactive=True)
+
+ src_image_interactive = image_processor_ins.system_para.get(
+ 'SRC_IMAGE_INTERACTIVE', True)
+ src_image_tool = image_processor_ins.system_para.get(
+ 'SRC_IMAGE_TOOL', preview_src_image_tool)
+ ret_src_image = gr.Image(interactive=src_image_interactive)
+
+ src_image_mask_interactive = image_processor_ins.system_para.get(
+ 'SRC_IMAGE_MASK_INTERACTIVE', True)
+ src_image_mask_tool = image_processor_ins.system_para.get(
+ 'SRC_IMAGE_MASK_TOOL', preview_src_mask_tool)
+ ret_src_mask = gr.Image(interactive=src_image_mask_interactive)
+ target_image_interactive = image_processor_ins.system_para.get(
+ 'TARGET_IMAGE_INTERACTIVE', True)
+ target_image_tool = image_processor_ins.system_para.get(
+ 'TARGET_IMAGE_TOOL', preview_target_image_tool)
+ ret_target_image = gr.Image(interactive=target_image_interactive)
+
+ caption_interactive = image_processor_ins.system_para.get(
+ 'CAPTION_INTERACTIVE', True)
+ ret_caption = gr.Textbox(interactive=caption_interactive)
+
+ preview_btn_visible = image_processor_ins.system_para.get(
+ 'PREVIEW_BTN_VISIBLE', True)
+ ret_preview_btn = gr.Column(visible=preview_btn_visible)
+
+ return (ret_height_ratio, ret_width_ratio, use_mask, ret_src_image,
+ ret_src_mask, ret_target_image, ret_caption,
+ ret_preview_btn, src_image_tool, src_image_mask_tool,
+ target_image_tool)
self.image_preprocess_method.change(
image_preprocess_method_change,
- inputs=[self.image_preprocess_method],
- outputs=[self.height_ratio, self.width_ratio],
+ inputs=[
+ self.image_preprocess_method, self.preview_src_image_tool,
+ self.preview_src_mask_image_tool, self.preview_taget_image_tool
+ ],
+ outputs=[
+ self.height_ratio, self.width_ratio, self.use_mask,
+ self.preview_src_image, self.preview_src_mask_image,
+ self.preview_taget_image, self.preview_caption,
+ self.image_preview_btn_panel, self.preview_src_image_tool,
+ self.preview_src_mask_image_tool, self.preview_taget_image_tool
+ ],
queue=False)
def caption_preprocess_method_change(caption_preprocess_method):
@@ -1253,10 +1831,10 @@ class DatasetGalleryUI(UIBase):
def preprocess_caption(mode_state, preprocess_method, sys_prompt,
max_new_tokens, min_new_tokens, num_beams,
- repetition_penalty, temperature,
+ repetition_penalty, temperature, use_local,
caption_update_mode, upload_image,
- upload_src_image, upload_caption, dataset_type,
- dataset_name):
+ upload_src_image, upload_src_mask,
+ upload_caption, dataset_type, dataset_name):
reverse_update_mode = {
v: idx
@@ -1290,34 +1868,47 @@ class DatasetGalleryUI(UIBase):
one_data = dataset_ins.data[index]
relative_image_path = one_data.get(
'edit_relative_path', one_data['relative_path'])
- src_image = os.path.join(
- dataset_ins.meta['local_work_dir'],
- relative_image_path)
- response = processor_ins(
- src_image,
- prompt=sys_prompt,
- max_new_tokens=max_new_tokens,
- min_new_tokens=min_new_tokens,
- num_beams=num_beams,
- repetition_penalty=repetition_penalty,
- temperature=temperature)
+ target_image = Image.open(
+ os.path.join(dataset_ins.meta['local_work_dir'],
+ relative_image_path))
if dataset_type == 'scepter_img2img':
relative_image_path = one_data.get(
'edit_relative_path',
one_data['src_relative_path'])
- src_image = os.path.join(
- dataset_ins.meta['local_work_dir'],
- relative_image_path)
- src_response = processor_ins(
- src_image,
- prompt=sys_prompt,
- max_new_tokens=max_new_tokens,
- min_new_tokens=min_new_tokens,
- num_beams=num_beams,
- repetition_penalty=repetition_penalty,
- temperature=temperature)
- response = 'src: ' + src_response + ' target: ' + response
+ src_image = Image.open(
+ os.path.join(dataset_ins.meta['local_work_dir'],
+ relative_image_path))
+ relative_src_mask_path = one_data.get(
+ 'edit_relative_path',
+ one_data['src_mask_relative_path'])
+ src_mask_image = Image.open(
+ os.path.join(dataset_ins.meta['local_work_dir'],
+ relative_src_mask_path))
+ else:
+ src_image = None
+ src_mask_image = None
+
+ kwargs = {
+ 'src_image': src_image,
+ 'src_mask': src_mask_image,
+ 'target_image': target_image,
+ 'caption': one_data.get('edit_caption', ''),
+ 'preview_src_image': None,
+ 'preview_src_mask': None,
+ 'preview_target_image': None,
+ 'preview_caption': None,
+ 'use_preview': False,
+ 'use_local': use_local,
+ 'sys_prompt': sys_prompt,
+ 'max_new_tokens': max_new_tokens,
+ 'min_new_tokens': min_new_tokens,
+ 'num_beams': num_beams,
+ 'repetition_penalty': repetition_penalty,
+ 'temperature': temperature,
+ 'cache': self.cache
+ }
+ response = processor_ins(**kwargs)
if update_mode == 0:
if len(dataset_ins.data[index]['edit_caption']) > 0:
@@ -1333,45 +1924,47 @@ class DatasetGalleryUI(UIBase):
visible=True)
ret_upload_caption = gr.Textbox()
elif mode_state == 'add':
- save_folder = os.path.join(dataset_ins.local_work_dir, 'cache')
- os.makedirs(save_folder, exist_ok=True)
if isinstance(upload_image, dict):
image = upload_image['image']
else:
image = upload_image
w, h = image.size
image_path = os.path.join(
- save_folder, f'{imagehash.phash(image)}_{w}_{h}.png')
+ self.cache, f'{imagehash.phash(image)}_{w}_{h}.jpg')
if not os.path.exists(image_path):
image.save(image_path)
- response = processor_ins(image_path,
- prompt=sys_prompt,
- max_new_tokens=max_new_tokens,
- min_new_tokens=min_new_tokens,
- num_beams=num_beams,
- repetition_penalty=repetition_penalty,
- temperature=temperature)
-
if dataset_type == 'scepter_img2img':
if isinstance(upload_src_image, dict):
src_image = upload_src_image['image']
else:
src_image = upload_src_image
- w, h = src_image.size
- src_image_path = os.path.join(
- save_folder,
- f'{imagehash.phash(src_image)}_{w}_{h}.png')
- if not os.path.exists(src_image_path):
- src_image.save(src_image_path)
- src_response = processor_ins(
- src_image_path,
- prompt=sys_prompt,
- max_new_tokens=max_new_tokens,
- min_new_tokens=min_new_tokens,
- num_beams=num_beams,
- repetition_penalty=repetition_penalty,
- temperature=temperature)
- response = 'src: ' + src_response + ' target: ' + response
+ if isinstance(upload_src_mask, dict):
+ src_mask = upload_src_mask['image']
+ else:
+ src_mask = upload_src_mask
+ else:
+ src_image = None
+ src_mask = None
+ kwargs = {
+ 'src_image': src_image,
+ 'src_mask': src_mask,
+ 'target_image': image_path,
+ 'caption': upload_caption,
+ 'preview_src_image': None,
+ 'preview_src_mask': None,
+ 'preview_target_image': None,
+ 'preview_caption': None,
+ 'use_preview': False,
+ 'use_local': use_local,
+ 'sys_prompt': sys_prompt,
+ 'max_new_tokens': max_new_tokens,
+ 'min_new_tokens': min_new_tokens,
+ 'num_beams': num_beams,
+ 'repetition_penalty': repetition_penalty,
+ 'temperature': temperature,
+ 'cache': self.cache
+ }
+ response = processor_ins(**kwargs)
ret_edit_caption = gr.Textbox()
if update_mode == 0:
@@ -1401,13 +1994,170 @@ class DatasetGalleryUI(UIBase):
self.mode_state, self.caption_preprocess_method,
self.sys_prompt, self.max_new_tokens, self.min_new_tokens,
self.num_beams, self.repetition_penalty, self.temperature,
- self.caption_update_mode, self.upload_image,
- self.upload_src_image, self.upload_caption,
- create_dataset.dataset_type, create_dataset.dataset_name
+ self.use_local, self.caption_update_mode, self.upload_image,
+ self.upload_src_image, self.upload_src_mask,
+ self.upload_caption, create_dataset.dataset_type,
+ create_dataset.dataset_name
],
outputs=[self.edit_caption, self.upload_caption, self.sys_log],
queue=False)
+ def preprocess_image_preview(preprocess_method, dataset_type,
+ src_image, src_mask, target_image,
+ caption, use_mask, height_ratio,
+ width_ratio):
+ dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
+ processor_ins = self.processors_manager.get_processor(
+ 'image', preprocess_method)
+ if processor_ins is None:
+ sys_log = 'Current processor is illegal'
+ return (gr.Image(), gr.Image(), gr.Image(),
+ self.component_names.system_log.format(sys_log))
+ is_flag, msg = processor_ins.load_model()
+ if not is_flag:
+ sys_log = f'Load processor failed: {msg}'
+ return (gr.Image(), gr.Image(), gr.Image(),
+ self.component_names.system_log.format(sys_log))
+ if not dataset_type == 'scepter_img2img':
+ src_image = None
+ src_mask = None
+ kwargs = {
+ 'width_ratio': width_ratio,
+ 'height_ratio': height_ratio,
+ 'use_mask': use_mask,
+ 'src_image': src_image,
+ 'src_mask': src_mask,
+ 'target_image': target_image,
+ 'caption': caption,
+ 'preview_target_image': target_image,
+ 'preview_src_image': src_image,
+ 'preview_src_mask': src_mask,
+ 'preview_caption': caption
+ }
+ processor_ins.load_model()
+ output_data = processor_ins(**kwargs)
+ processor_ins.unload_model()
+ if output_data.get('target_image', None) is not None:
+ target_image = gr.Image(value=output_data['target_image'])
+ else:
+ target_image = gr.Image()
+ if output_data.get('src_image', None) is not None:
+ src_image = gr.Image(value=output_data.get('src_image', None))
+ else:
+ src_image = gr.Image()
+ if output_data.get('src_mask', None) is not None:
+ src_mask = gr.Image(value=output_data.get('src_mask', None))
+ else:
+ src_mask = gr.Image()
+ return (src_image, src_mask, target_image, '')
+
+ self.image_preview_btn.click(
+ preprocess_image_preview,
+ inputs=[
+ self.image_preprocess_method, create_dataset.dataset_type,
+ self.preview_src_image, self.preview_src_mask_image,
+ self.preview_taget_image, self.preview_caption, self.use_mask,
+ self.height_ratio, self.width_ratio
+ ],
+ outputs=[
+ self.preview_src_image, self.preview_src_mask_image,
+ self.preview_taget_image, self.sys_log
+ ])
+
+ def preprocess_caption_preview(
+ preprocess_method, dataset_type, sys_prompt, max_new_tokens,
+ min_new_tokens, num_beams, repetition_penalty, temperature,
+ use_local, caption_update_mode, preview_src_image,
+ preview_src_mask_image, preview_target_image, preview_caption):
+ reverse_update_mode = {
+ v: idx
+ for idx, v in enumerate(
+ self.component_names.caption_update_choices)
+ }
+
+ update_mode = reverse_update_mode.get(caption_update_mode, -1)
+
+ processor_ins = self.processors_manager.get_processor(
+ 'caption', preprocess_method)
+ if processor_ins is None:
+ sys_log = 'Current processor is illegal'
+ return gr.Textbox(), gr.Textbox(
+ ), self.component_names.system_log.format(sys_log)
+
+ is_flag, msg = processor_ins.load_model()
+ if not is_flag:
+ sys_log = f'Load processor failed: {msg}'
+ return gr.Textbox(), self.component_names.system_log.format(
+ sys_log)
+ dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
+
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['background']
+ else:
+ prev_target_image = preview_target_image
+
+ if dataset_type == 'scepter_img2img':
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background']
+ else:
+ prev_src_image = preview_src_image
+
+ if isinstance(preview_src_mask_image, dict):
+ prev_src_mask = preview_src_mask_image['layers'][0]
+ else:
+ prev_src_mask = preview_src_mask_image
+
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ kwargs = {
+ 'src_image': None,
+ 'src_mask': None,
+ 'target_image': None,
+ 'caption': None,
+ 'preview_src_image': prev_src_image,
+ 'preview_src_mask': prev_src_mask,
+ 'preview_target_image': prev_target_image,
+ 'preview_caption': preview_caption,
+ 'use_preview': True,
+ 'use_local': use_local,
+ 'sys_prompt': sys_prompt,
+ 'max_new_tokens': max_new_tokens,
+ 'min_new_tokens': min_new_tokens,
+ 'num_beams': num_beams,
+ 'repetition_penalty': repetition_penalty,
+ 'temperature': temperature,
+ 'cache': self.cache
+ }
+ response = processor_ins(**kwargs)
+
+ if update_mode == 0:
+ if len(preview_caption) > 0:
+ preview_caption += ';'
+ preview_caption += response
+ elif update_mode == 1:
+ preview_caption = response
+
+ is_flag, msg = processor_ins.unload_model()
+ if not is_flag:
+ sys_log = f'Unoad processor failed: {msg}'
+ return gr.Textbox(), self.component_names.system_log.format(
+ sys_log)
+ return gr.Textbox(value=preview_caption), ''
+
+ self.caption_preview_btn.click(
+ preprocess_caption_preview,
+ inputs=[
+ self.caption_preprocess_method, create_dataset.dataset_type,
+ self.sys_prompt, self.max_new_tokens, self.min_new_tokens,
+ self.num_beams, self.repetition_penalty, self.temperature,
+ self.use_local, self.caption_update_mode,
+ self.preview_src_image, self.preview_src_mask_image,
+ self.preview_taget_image, self.preview_caption
+ ],
+ outputs=[self.preview_caption, self.sys_log])
+
def mode_state_change(mode_state, dataset_type, dataset_name):
# default is editing current sample
dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
@@ -1415,32 +2165,16 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None or len(dataset_ins) < 1:
- return (
- gr.Gallery(),
- gr.Gallery(),
- gr.Column(),
- gr.Row(visible=True),
- gr.Column(),
- gr.Column(),
- gr.Image(),
- gr.Column(),
- '',
- gr.Markdown(),
- gr.Textbox(),
- gr.CheckboxGroup(),
- gr.Column(),
- # gr.Column(),
- gr.Row(),
- # gr.Column(),
- gr.Row(),
- gr.Column(),
- gr.Gallery(),
- gr.Gallery(),
- gr.Column(),
- gr.Dropdown(),
- gr.Dropdown(value=[]),
- self.component_names.system_log.format(
- self.component_names.illegal_blank_dataset))
+ return (gr.Gallery(), gr.Gallery(),
+ gr.Image(), gr.Column(), gr.Row(visible=True),
+ gr.Column(), gr.Column(), gr.Image(), gr.Row(), '',
+ gr.Markdown(), gr.Textbox(), gr.CheckboxGroup(),
+ gr.Column(), gr.Row(), gr.Row(), gr.Column(),
+ gr.Gallery(), gr.Textbox(), gr.Gallery(),
+ gr.Image(), gr.Column(), gr.Dropdown(),
+ gr.Dropdown(value=[]),
+ self.component_names.system_log.format(
+ self.component_names.illegal_blank_dataset))
dataset_ins.set_edit_range(str(dataset_ins.cursor + 1))
image_list = [
os.path.join(
@@ -1476,103 +2210,65 @@ class DatasetGalleryUI(UIBase):
selected_index=None,
visible=True)
ret_edit_src_panel = gr.Column(visible=True)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ edit_mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info.get(
+ 'edit_src_mask_relative_path',
+ current_info['src_mask_relative_path']))
+ else:
+ edit_mask_path = None
+ ret_edit_mask = gr.Image(value=edit_mask_path,
+ visible=True)
else:
ret_edit_src_gallery = gr.Gallery(visible=False)
ret_edit_src_panel = gr.Column(visible=False)
- return (
- gr.Gallery(),
- gr.Gallery(),
- gr.Column(),
- gr.Row(visible=True),
- gr.Column(visible=True),
- gr.Column(visible=False),
- gr.Image(),
- gr.Column(),
- '',
- gr.Markdown(visible=False),
- gr.Textbox(),
- gr.CheckboxGroup(value=None),
- gr.Column(visible=False),
- # gr.Column(visible=False),
- gr.Row(visible=False),
- # gr.Column(visible=False),
- gr.Row(visible=False),
- gr.Column(visible=True),
- gr.Gallery(value=image_list,
- selected_index=selected_index,
- visible=True),
- ret_edit_caption,
- ret_edit_src_gallery,
- ret_edit_src_panel,
- gr.Dropdown(
- # choices=self.component_names.range_mode_name,
- # value=self.component_names.range_mode_name[0]
- ),
- gr.Dropdown(value=[]),
- self.component_names.system_log.format(''))
+ ret_edit_mask = gr.Image(visible=False)
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Column(),
+ gr.Row(visible=True), gr.Column(visible=True),
+ gr.Column(visible=False), gr.Image(), gr.Row(), '',
+ gr.Markdown(visible=False), gr.Textbox(),
+ gr.CheckboxGroup(value=None), gr.Column(visible=False),
+ gr.Row(visible=False), gr.Row(visible=False),
+ gr.Column(visible=True),
+ gr.Gallery(value=image_list,
+ selected_index=selected_index,
+ visible=True), ret_edit_caption,
+ ret_edit_src_gallery,
+ ret_edit_mask, ret_edit_src_panel, gr.Dropdown(),
+ gr.Dropdown(value=[]),
+ self.component_names.system_log.format(''))
elif mode_state == 'add':
- return (
- gr.Gallery(),
- gr.Gallery(),
- gr.Column(),
- gr.Row(visible=False),
- gr.Column(visible=False),
- gr.Column(visible=True),
- gr.Image(),
- gr.Column(visible=True) if dataset_type
- == 'scepter_img2img' else gr.Column(visible=False),
- '',
- gr.Markdown(visible=True) if dataset_type
- == 'scepter_img2img' else gr.Markdown(visible=False),
- gr.Textbox(value=''),
- gr.CheckboxGroup(),
- gr.Column(visible=True),
- # gr.Column(visible=True),
- gr.Row(visible=True),
- # gr.Column(visible=True),
- gr.Row(visible=True),
- gr.Column(visible=False),
- gr.Gallery(visible=False),
- gr.Textbox(),
- gr.Gallery(visible=False),
- gr.Column(visible=False),
- gr.Dropdown(
- # choices=self.component_names.range_mode_name,
- # value=self.component_names.range_mode_name[0]
- ),
- gr.Dropdown(value=[]),
- self.component_names.system_log.format(''))
+ return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Column(),
+ gr.Row(visible=False), gr.Column(visible=False),
+ gr.Column(visible=True), gr.Image(),
+ gr.Row(visible=True) if dataset_type
+ == 'scepter_img2img' else gr.Column(visible=False), '',
+ gr.Markdown(visible=True) if dataset_type
+ == 'scepter_img2img' else gr.Markdown(visible=False),
+ gr.Textbox(value=''), gr.CheckboxGroup(),
+ gr.Column(visible=False), gr.Row(visible=False),
+ gr.Row(visible=False), gr.Column(visible=False),
+ gr.Gallery(visible=False), gr.Textbox(),
+ gr.Gallery(visible=False), gr.Image(visible=False),
+ gr.Column(visible=False), gr.Dropdown(),
+ gr.Dropdown(value=[]),
+ self.component_names.system_log.format(''))
else:
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return (
- gr.Gallery(),
- gr.Gallery(),
- gr.Column(),
- gr.Row(visible=True),
- gr.Column(),
- gr.Column(),
- gr.Image(),
- gr.Column(),
- '',
- gr.Markdown(),
- gr.Textbox(),
- gr.CheckboxGroup(),
- gr.Column(),
- # gr.Column(),
- gr.Row(),
- # gr.Column(),
- gr.Row(),
- gr.Column(),
- gr.Gallery(),
- gr.Textbox(),
- gr.Gallery(),
- gr.Column(),
- gr.Dropdown(),
- gr.Dropdown(value=[]),
- self.component_names.system_log.format(
- self.component_names.illegal_blank_dataset))
+ return (gr.Gallery(), gr.Gallery(),
+ gr.Image(), gr.Column(), gr.Row(visible=True),
+ gr.Column(), gr.Column(), gr.Image(), gr.Row(), '',
+ gr.Markdown(), gr.Textbox(), gr.CheckboxGroup(),
+ gr.Column(), gr.Row(), gr.Row(), gr.Column(),
+ gr.Gallery(), gr.Textbox(), gr.Gallery(),
+ gr.Image(), gr.Column(), gr.Dropdown(),
+ gr.Dropdown(value=[]),
+ self.component_names.system_log.format(
+ self.component_names.illegal_blank_dataset))
cursor = dataset_ins.cursor
image_list = [
os.path.join(dataset_ins.meta['local_work_dir'],
@@ -1595,9 +2291,28 @@ class DatasetGalleryUI(UIBase):
selected_index=None,
visible=True)
ret_src_panel = gr.Column(visible=True)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['src_mask_relative_path'])
+ else:
+ mask_path = None
+ ret_mask = gr.Image(value=mask_path, visible=True)
+ if len(current_info) > 0:
+ edit_mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info.get(
+ 'edit_src_mask_relative_path',
+ current_info['src_mask_relative_path']))
+ else:
+ edit_mask_path = None
+ edit_mask = gr.Image(value=edit_mask_path, visible=True)
else:
ret_src_gallery = gr.Gallery(visible=False)
+ ret_mask = gr.Image()
ret_src_panel = gr.Column(visible=False)
+ edit_mask = gr.Image()
if cursor >= 0:
ret_gl_gallery = gr.Gallery(label=dataset_name,
@@ -1608,35 +2323,17 @@ class DatasetGalleryUI(UIBase):
value=image_list,
selected_index=None)
- return (
- ret_gl_gallery,
- ret_src_gallery,
- ret_src_panel,
- gr.Row(visible=False),
- gr.Column(visible=False),
- gr.Column(visible=False),
- gr.Image(),
- gr.Column(),
- '',
- gr.Markdown(visible=False),
- gr.Textbox(value=''),
- gr.CheckboxGroup(value=None),
- gr.Column(visible=False),
- # gr.Column(visible=False),
- gr.Row(visible=False),
- # gr.Column(visible=False),
- gr.Row(visible=False),
- gr.Column(visible=False),
- gr.Gallery(visible=False),
- gr.Textbox(),
- gr.Gallery(visible=False),
- gr.Column(visible=False),
- gr.Dropdown(
- # choices=self.component_names.range_mode_name,
- # value=self.component_names.range_mode_name[0]
- ),
- gr.Dropdown(value=[]),
- '')
+ return (ret_gl_gallery, ret_src_gallery, ret_mask,
+ ret_src_panel, gr.Row(visible=False),
+ gr.Column(visible=False), gr.Column(visible=False),
+ gr.Image(), gr.Row(), '', gr.Markdown(visible=False),
+ gr.Textbox(value=''), gr.CheckboxGroup(value=None),
+ gr.Column(visible=False), gr.Row(visible=False),
+ gr.Row(visible=False), gr.Column(visible=False),
+ gr.Gallery(visible=False), gr.Textbox(),
+ gr.Gallery(visible=False), edit_mask,
+ gr.Column(visible=False), gr.Dropdown(),
+ gr.Dropdown(value=[]), '')
self.mode_state.change(
mode_state_change,
@@ -1645,16 +2342,55 @@ class DatasetGalleryUI(UIBase):
create_dataset.dataset_name
],
outputs=[
- self.gl_dataset_images, self.src_gl_dataset_images,
- self.src_panel, self.edit_setting_panel,
- self.edit_confirm_panel, self.upload_panel, self.upload_image,
- self.upload_src_image_panel, self.upload_image_info,
- self.upload_src_image_info, self.upload_caption,
- self.preprocess_checkbox, self.preprocess_panel,
- self.image_preprocess_panel, self.caption_preprocess_panel,
- self.edit_panel, self.edit_gl_dataset_images,
- self.edit_caption, self.edit_src_gl_dataset_images,
- self.edit_src_panel, self.range_mode, self.data_range,
+ # gr.Gallery
+ self.gl_dataset_images,
+ # gr.Gallery
+ self.src_gl_dataset_images,
+ # gr.Image
+ self.src_mask,
+ # gr.Column
+ self.src_panel,
+ # gr.Row
+ self.edit_setting_panel,
+ # gr.Column
+ self.edit_confirm_panel,
+ # gr.Column
+ self.upload_panel,
+ # gr.Image
+ self.upload_image,
+ # gr.Row
+ self.upload_src_image_panel,
+ # gr.Markdown
+ self.upload_image_info,
+ # gr.Markdown
+ self.upload_src_image_info,
+ # gr.Textbox
+ self.upload_caption,
+ # gr.CheckboxGroup
+ self.preprocess_checkbox,
+ # gr.Column
+ self.preprocess_panel,
+ # gr.Row
+ self.image_preprocess_panel,
+ # gr.Row
+ self.caption_preprocess_panel,
+ # gr.Column
+ self.edit_panel,
+ # gr.Gallery
+ self.edit_gl_dataset_images,
+ # gr.Textbox
+ self.edit_caption,
+ # gr.Gallery
+ self.edit_src_gl_dataset_images,
+ # gr.Image
+ self.edit_src_mask,
+ # gr.Column
+ self.edit_src_panel,
+ # gr.Dropdown
+ self.range_mode,
+ # gr.Dropdown
+ self.data_range,
+ # gr.Markdown
self.sys_log
],
queue=False)
@@ -1662,14 +2398,16 @@ class DatasetGalleryUI(UIBase):
def range_change(range_mode, dataset_type, dataset_name):
hit_range_mode = range_state_trans(range_mode)
if hit_range_mode == 2:
- return (gr.Dropdown(visible=True), gr.Gallery(), gr.Gallery())
+ return (gr.Dropdown(visible=True), gr.Gallery(), gr.Gallery(),
+ gr.Image())
else:
dataset_type = create_dataset.get_trans_dataset_type(
dataset_type)
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return (gr.Dropdown(), gr.Gallery(), gr.Gallery())
+ return (gr.Dropdown(), gr.Gallery(), gr.Gallery(),
+ gr.Image())
if hit_range_mode == 1:
dataset_ins.set_edit_range(-1)
else:
@@ -1700,12 +2438,24 @@ class DatasetGalleryUI(UIBase):
ret_edit_src_gallery = gr.Gallery(label=dataset_name,
value=src_image_list,
selected_index=None)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ edit_mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info.get(
+ 'edit_src_mask_relative_path',
+ current_info['src_mask_relative_path']))
+ else:
+ edit_mask_path = None
+ edit_mask = gr.Image(value=edit_mask_path)
else:
ret_edit_src_gallery = gr.Gallery()
+ edit_mask = gr.Image()
return (gr.Dropdown(visible=False),
gr.Gallery(value=image_list,
selected_index=selected_index,
- visible=True), ret_edit_src_gallery)
+ visible=True), ret_edit_src_gallery,
+ edit_mask)
self.range_mode.change(range_change,
inputs=[
@@ -1716,27 +2466,30 @@ class DatasetGalleryUI(UIBase):
outputs=[
self.data_range,
self.edit_gl_dataset_images,
- self.edit_src_gl_dataset_images
+ self.edit_src_gl_dataset_images,
+ self.edit_src_mask
],
queue=False)
def set_range(data_range, dataset_type, dataset_name):
if len(data_range) == 0:
- return (gr.Gallery(), gr.Gallery(), gr.Gallery(), gr.Textbox(),
- gr.Text(), gr.Textbox(), gr.Text(), gr.Text(),
- self.component_names.system_log.format(''))
+ return (gr.Gallery(), gr.Gallery(), gr.Gallery(), gr.Image(),
+ gr.Textbox(), gr.Text(), gr.Textbox(), gr.Text(),
+ gr.Text(), self.component_names.system_log.format(''))
data_range = ','.join(data_range)
dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return (gr.Gallery(), gr.Gallery(), gr.Gallery(), gr.Textbox(),
- gr.Text(), gr.Textbox(), gr.Text(), gr.Text(), 'None')
+ return (gr.Gallery(), gr.Gallery(), gr.Gallery(), gr.Image(),
+ gr.Textbox(), gr.Text(), gr.Textbox(), gr.Text(),
+ gr.Text(), 'None')
flg, msg = dataset_ins.set_edit_range(data_range)
if not flg:
sys_log = self.component_names.system_log.format(msg)
- return (gr.Gallery(), gr.Gallery(), gr.Gallery(), gr.Textbox(),
- gr.Text(), gr.Textbox(), gr.Text(), gr.Text(), sys_log)
+ return (gr.Gallery(), gr.Gallery(), gr.Gallery(), gr.Image(),
+ gr.Textbox(), gr.Text(), gr.Textbox(), gr.Text(),
+ gr.Text(), sys_log)
image_list = [
os.path.join(dataset_ins.meta['local_work_dir'],
v.get('edit_relative_path', v['relative_path']))
@@ -1767,8 +2520,19 @@ class DatasetGalleryUI(UIBase):
value=src_image_list,
selected_index=selected_index,
visible=True)
+ if len(current_info) > 0:
+ edit_mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info.get(
+ 'edit_src_mask_relative_path',
+ current_info['src_mask_relative_path']))
+ else:
+ edit_mask_path = None
+ ret_edit_src_mask = gr.Image(value=edit_mask_path,
+ visible=True)
else:
ret_edit_src_image_gl = gr.Gallery()
+ ret_edit_src_mask = gr.Image()
ret_caption = gr.Textbox(value=current_info.get('caption', ''),
visible=True)
@@ -1803,6 +2567,7 @@ class DatasetGalleryUI(UIBase):
ret_image_gl = gr.Gallery()
ret_edit_image_gl = gr.Gallery()
ret_edit_src_image_gl = gr.Gallery()
+ ret_edit_src_mask = gr.Image()
ret_caption = gr.Textbox()
ret_edit_caption = gr.Textbox()
@@ -1810,8 +2575,9 @@ class DatasetGalleryUI(UIBase):
edit_image_info = ''
ret_info = gr.Text()
return (ret_image_gl, ret_edit_image_gl, ret_edit_src_image_gl,
- ret_caption, image_info, ret_edit_caption, edit_image_info,
- ret_info, self.component_names.system_log.format(''))
+ ret_edit_src_mask, ret_caption, image_info,
+ ret_edit_caption, edit_image_info, ret_info,
+ self.component_names.system_log.format(''))
self.data_range.change(
set_range,
@@ -1821,9 +2587,9 @@ class DatasetGalleryUI(UIBase):
],
outputs=[
self.gl_dataset_images, self.edit_gl_dataset_images,
- self.edit_src_gl_dataset_images, self.ori_caption,
- self.image_info, self.edit_caption, self.edit_image_info,
- self.info, self.sys_log
+ self.edit_src_gl_dataset_images, self.edit_src_mask,
+ self.ori_caption, self.image_info, self.edit_caption,
+ self.edit_image_info, self.info, self.sys_log
],
queue=False)
@@ -1867,7 +2633,7 @@ class DatasetGalleryUI(UIBase):
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
if dataset_ins is None:
- return (gr.Gallery(), gr.Row(visible=True),
+ return (gr.Gallery(), gr.Image(), gr.Row(visible=True),
self.component_names.system_log.format(
self.component_names.illegal_blank_dataset))
if src_dataset_state:
@@ -1886,11 +2652,22 @@ class DatasetGalleryUI(UIBase):
ret_src_image_gl = gr.Gallery(value=[],
selected_index=None,
visible=True)
+ current_info = dataset_ins.current_record
+ if len(current_info) > 0:
+ mask_path = os.path.join(
+ dataset_ins.meta['local_work_dir'],
+ current_info['src_mask_relative_path'])
+ else:
+ mask_path = None
+ ret_mask = gr.Image(value=mask_path, visible=True)
else:
ret_src_image_gl = gr.Gallery(visible=False)
+ ret_mask = gr.Image(visible=False)
else:
ret_src_image_gl = gr.Gallery()
- return (ret_src_image_gl, gr.Row(visible=True),
+ ret_mask = gr.Image()
+
+ return (ret_src_image_gl, ret_mask, gr.Row(visible=True),
self.component_names.system_log.format(''))
self.src_dataset_state.change(src_dataset_change,
@@ -1901,6 +2678,7 @@ class DatasetGalleryUI(UIBase):
],
outputs=[
self.src_gl_dataset_images,
+ self.src_mask,
self.edit_setting_panel, self.sys_log
],
queue=False)
@@ -1924,14 +2702,19 @@ class DatasetGalleryUI(UIBase):
else:
image = upload_image
w, h = image.size
- return gr.Markdown(
+ ret_mask = Image.new('L', (w, h), 0)
+ cache_path = os.path.join(self.cache,
+ f'int{time.time() * 100}.jpg')
+ ret_mask.save(cache_path)
+ return (gr.Markdown(
value=self.component_names.upload_src_image_info.format(h, w),
- visible=True)
+ visible=True), gr.Image(value=cache_path, visible=True))
- self.upload_src_image.upload(image_src_upload,
- inputs=[self.upload_src_image],
- outputs=[self.upload_src_image_info],
- queue=False)
+ self.upload_src_image.upload(
+ image_src_upload,
+ inputs=[self.upload_src_image],
+ outputs=[self.upload_src_image_info, self.upload_src_mask],
+ queue=False)
def image_clear():
return gr.Markdown(visible=False)
@@ -1959,51 +2742,68 @@ class DatasetGalleryUI(UIBase):
)
def add_file(dataset_type, dataset_name, upload_image,
- upload_src_image, caption):
- if upload_image is None:
- return (gr.Gallery(), gr.Textbox(), '', gr.Textbox(), '',
- gr.Text(), gr.Image(), gr.Image(), gr.Text(), '', '',
- gr.Text())
- if isinstance(upload_image, dict):
- image = upload_image['image']
- else:
- image = upload_image
-
- if isinstance(upload_src_image, dict):
- src_image = upload_src_image['image']
- else:
- src_image = upload_src_image
-
+ upload_src_image, upload_src_mask, caption):
dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
dataset_ins = create_dataset.dataset_dict.get(
dataset_type, {}).get(dataset_name, None)
+ # allow add blank image
+ if upload_image is not None:
+ if isinstance(upload_image, dict):
+ image = upload_image['image']
+ else:
+ image = upload_image
+ else:
+ image = None
+ if upload_src_image is not None:
+ if isinstance(upload_src_image, dict):
+ src_image = upload_src_image['image']
+ src_mask = upload_src_mask
+ else:
+ src_image = upload_src_image
+ src_mask = upload_src_mask
+ else:
+ src_image = None
+ src_mask = None
+ # if dataset_type == 'scepter_img2img', we allow the target image is blank image
+ if dataset_type == 'scepter_img2img' and src_image is not None and image is None:
+ w, h = src_image.size
+ image = Image.new('RGB', (w, h), (0, 0, 0))
+ if image is None:
+ return (gr.Image(value=None), gr.Image(value=None),
+ gr.Image(value=None), gr.Text(value=''), '', '',
+ gr.Text(value='view'))
+ if dataset_type == 'scepter_img2img' and (src_image is None
+ or src_mask is None):
+ return (gr.Image(value=None), gr.Image(value=None),
+ gr.Image(value=None), gr.Text(value=''), '', '',
+ gr.Text(value='view'))
if dataset_ins is None:
- return (gr.Gallery(), gr.Textbox(), '', gr.Textbox(), '',
- gr.Text(), gr.Image(), gr.Image(), gr.Text(), '', '',
- gr.Text())
- dataset_ins.add_record(image, caption, src_image=src_image)
+ return (gr.Image(value=None), gr.Image(value=None),
+ gr.Image(value=None), gr.Text(value=''), '', '',
+ gr.Text(value='view'))
+ dataset_ins.add_record(image,
+ caption,
+ src_image=src_image,
+ src_mask=src_mask)
return (gr.Image(value=None), gr.Image(value=None),
- gr.Text(value=''), '', '', gr.Text(value='view'))
+ gr.Image(value=None), gr.Text(value=''), '', '',
+ gr.Text(value='view'))
- self.upload_button.click(
- add_file,
- inputs=[
- create_dataset.dataset_type, create_dataset.dataset_name,
- self.upload_image, self.upload_src_image, self.upload_caption
- ],
- outputs=[
- # self.gl_dataset_images,
- # self.ori_caption,
- # self.image_info, self.edit_caption,
- # self.edit_image_info, self.info,
- self.upload_image,
- self.upload_src_image,
- self.upload_caption,
- self.upload_image_info,
- self.upload_src_image_info,
- self.mode_state
- ],
- queue=False)
+ self.upload_button.click(add_file,
+ inputs=[
+ create_dataset.dataset_type,
+ create_dataset.dataset_name,
+ self.upload_image, self.upload_src_image,
+ self.upload_src_mask, self.upload_caption
+ ],
+ outputs=[
+ self.upload_image, self.upload_src_image,
+ self.upload_src_mask, self.upload_caption,
+ self.upload_image_info,
+ self.upload_src_image_info,
+ self.mode_state
+ ],
+ queue=False)
def cancel_add_file():
return gr.Text(value='view')
diff --git a/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py b/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py
index aeeb59f..9911717 100644
--- a/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py
+++ b/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py
@@ -66,14 +66,14 @@ class ExportDatasetUI(UIBase):
gr.Textbox(
value=os.path.abspath(dataset_ins.local_work_dir)),
gr.Textbox(value=dataset_name))
-
- self.go_to_train.click(
- go_to_train,
- inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
- outputs=[
- manager.tabs, manager.self_train.trainer_ui.data_source,
- manager.self_train.trainer_ui.data_type,
- manager.self_train.trainer_ui.ms_data_name,
- manager.self_train.trainer_ui.ori_data_name
- ],
- queue=False)
+ if hasattr(manager, 'self_train'):
+ self.go_to_train.click(
+ go_to_train,
+ inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
+ outputs=[
+ manager.tabs, manager.self_train.trainer_ui.data_source,
+ manager.self_train.trainer_ui.data_type,
+ manager.self_train.trainer_ui.ms_data_name,
+ manager.self_train.trainer_ui.ori_data_name
+ ],
+ queue=False)
diff --git a/scepter/studio/preprocess/processors/caption_processors.py b/scepter/studio/preprocess/processors/caption_processors.py
index 4568252..b63f9e0 100644
--- a/scepter/studio/preprocess/processors/caption_processors.py
+++ b/scepter/studio/preprocess/processors/caption_processors.py
@@ -4,16 +4,30 @@
import numbers
import re
import time
-
+from torchvision.transforms.functional import InterpolationMode
+import torchvision.transforms as T
import torch
from PIL import Image
+from transformers import AutoModel, AutoTokenizer
+
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
+import numpy as np
from scepter.studio.preprocess.processors.base_processor import \
BaseCaptionProcessor
-__all__ = ['BlipImageBase', 'QWVL', 'QWVLQuantize']
+__all__ = ['BlipImageBase', 'QWVL', 'QWVLQuantize', 'InternVL15']
+def get_region(image, mask, mask_id):
+ locs = np.where(np.array(mask) == mask_id)
+ if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
+ return None
+ left, right = np.min(locs[1]), np.max(locs[1])
+ top, bottom = np.min(locs[0]), np.max(locs[0])
+ box = [left, top, right, bottom]
+ region_image = image.crop(box)
+ region_image.save("1.jpg")
+ return region_image
class BlipImageBase(BaseCaptionProcessor):
def __init__(self, cfg, language='en'):
@@ -79,13 +93,40 @@ class BlipImageBase(BaseCaptionProcessor):
torch.cuda.ipc_collect()
return True, ''
- def __call__(self, image, prompt=None, **kwargs):
- raw_image = Image.open(image).convert('RGB')
- inputs = self.model_info['processor'](
- raw_image, return_tensors='pt').to(we.device_id)
+ def get_caption(self, image, use_local, mask):
+ image = image.convert('RGB')
+ if use_local:
+ image = get_region(image,
+ mask.convert('L'),
+ 255)
+
+ if image is None:
+ return ""
+ inputs = self.model_info['processor'](image, return_tensors='pt').to(we.device_id)
out = self.model_info['model'].generate(**inputs)
- return self.model_info['processor'].decode(out[0],
- skip_special_tokens=True)
+ caption = self.model_info['processor'].decode(out[0], skip_special_tokens=True)
+ return caption
+ def __call__(self, **kwargs):
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ use_local = kwargs.get('use_local', False)
+ cache = kwargs.get('cache', None)
+ preview_target_image = kwargs.get('preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.get('preview_src_mask', None) if use_preview else src_mask
+ preview_src_image = kwargs.get('preview_src_image', None) if use_preview else src_image
+ preview_caption = kwargs.get('preview_caption', None) if use_preview else caption
+ response = ""
+
+ if preview_src_image is not None:
+ response += ("src_caption" + ": " +
+ self.get_caption(preview_src_image, use_local, preview_src_mask) + "\n")
+ if preview_target_image is not None:
+ response += ("target_caption" + ": " +
+ self.get_caption(preview_target_image, use_local, preview_src_mask) + "\n")
+ return response
class QWVL(BaseCaptionProcessor):
@@ -158,11 +199,14 @@ class QWVL(BaseCaptionProcessor):
torch.cuda.ipc_collect()
return True, ''
- def __call__(self,
- image,
- prompt='Generate the caption in English',
- **kwargs):
-
+ def get_caption(self, image, prompt, kwargs, use_local, mask):
+ image = image.convert('RGB')
+ if use_local:
+ image = get_region(image,
+ mask.convert('L'),
+ 255)
+ if image is None:
+ return ""
torch.manual_seed(int(time.time()) % 100000)
query = self.model_info['tokenizer'].from_list_format([
{
@@ -172,8 +216,6 @@ class QWVL(BaseCaptionProcessor):
'text': prompt
},
])
- print(kwargs)
-
inputs = self.model_info['tokenizer'](query, return_tensors='pt')
inputs = inputs.to(self.model_info['device'])
pred = self.model_info['model'].generate(**inputs, **kwargs)
@@ -189,6 +231,32 @@ class QWVL(BaseCaptionProcessor):
regex = re.compile(r'^[\-\_]+')
ret_caption = re.sub(regex, r'', ret_caption)
return ret_caption
+ def __call__(self, **kwargs):
+ target_image = kwargs.pop('target_image', None)
+ src_mask = kwargs.pop('src_mask', None)
+ src_image = kwargs.pop('src_image', None)
+ use_preview = kwargs.pop('use_preview', True)
+ sys_prompt = kwargs.pop('sys_prompt', 'Generate the caption in English')
+ caption = kwargs.pop('caption', None)
+ use_local = kwargs.pop('use_local', False)
+ cache = kwargs.pop('cache', None)
+ preview_target_image = kwargs.pop('preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.pop('preview_src_mask', None) if use_preview else src_mask
+ preview_src_image = kwargs.pop('preview_src_image', None) if use_preview else src_image
+ preview_caption = kwargs.pop('preview_caption', None) if use_preview else caption
+ kwargs_keys = ["max_new_tokens", "min_new_tokens", "num_beams", "repetition_penalty", "temperature"]
+ process_kwargs = {}
+ for key in kwargs_keys:
+ if key in kwargs:
+ process_kwargs[key] = kwargs[key]
+ response = ""
+ if preview_src_image is not None:
+ response += ("src_caption" + ": " +
+ self.get_caption(preview_src_image, sys_prompt, process_kwargs, use_local, preview_src_mask) + "\n")
+ if preview_target_image is not None:
+ response += ("target_caption" + ": " +
+ self.get_caption(preview_target_image, sys_prompt, process_kwargs, use_local, preview_src_mask) + "\n")
+ return response
class QWVLQuantize(QWVL):
@@ -255,3 +323,147 @@ class QWVLQuantize(QWVL):
self.model_info['device'] = 'offline'
torch.cuda.empty_cache()
return True, ''
+
+
+class InternVL15(QWVL):
+ # AI-ModelScope/InternVL-Chat-V1-5
+ def load_model(self):
+ is_flg, msg = super(QWVL, self).load_model()
+ if not is_flg:
+ return is_flg, msg
+ self.model = None
+ if self.model_info['device'] == 'offline':
+ try:
+ local_path = FS.get_dir_to_local_dir(self.model_path)
+ # If you have an 80G A100 GPU, you can put the entire model on a single GPU.
+ model = AutoModel.from_pretrained(
+ local_path,
+ torch_dtype=torch.bfloat16,
+ low_cpu_mem_usage=True,
+ trust_remote_code=True).eval().to(we.device_id)
+ tokenizer = AutoTokenizer.from_pretrained(local_path, trust_remote_code=True)
+ # model.to(we.device_id)
+ except Exception as e:
+ if self.model is not None:
+ del self.model
+ return False, f"Load model error '{e}'"
+ self.model_info['device'] = model.device
+ self.model_info['model'] = model
+ self.model_info['tokenizer'] = tokenizer
+ elif self.model_info['device'] == 'cpu':
+ try:
+ self.model_info['model'].to(we.device_id)
+ self.model_info['device'] = we.device_id
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+ except Exception as e:
+ del self.model_info['model']
+ self.model_info['model'] = None
+ self.model_info['device'] = 'offline'
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+ return False, f"Load model error '{e}'"
+
+ return True, ''
+
+ def unload_model(self):
+ print(self.model_info['device'])
+ if (isinstance(self.model_info['device'], numbers.Number)
+ or str(self.model_info['device']).startswith('cuda')):
+ del self.model_info['model']
+ self.model_info['model'] = None
+ self.model_info['device'] = 'offline'
+ torch.cuda.empty_cache()
+ return True, ''
+ def build_transform(self, input_size):
+ IMAGENET_MEAN = (0.485, 0.456, 0.406)
+ IMAGENET_STD = (0.229, 0.224, 0.225)
+ MEAN, STD = IMAGENET_MEAN, IMAGENET_STD
+ transform = T.Compose([
+ T.Lambda(lambda img: img.convert('RGB') if img.mode != 'RGB' else img),
+ T.Resize((input_size, input_size), interpolation=InterpolationMode.BICUBIC),
+ T.ToTensor(),
+ T.Normalize(mean=MEAN, std=STD)
+ ])
+ return transform
+
+
+ def find_closest_aspect_ratio(self, aspect_ratio, target_ratios, width, height, image_size):
+ best_ratio_diff = float('inf')
+ best_ratio = (1, 1)
+ area = width * height
+ for ratio in target_ratios:
+ target_aspect_ratio = ratio[0] / ratio[1]
+ ratio_diff = abs(aspect_ratio - target_aspect_ratio)
+ if ratio_diff < best_ratio_diff:
+ best_ratio_diff = ratio_diff
+ best_ratio = ratio
+ elif ratio_diff == best_ratio_diff:
+ if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:
+ best_ratio = ratio
+ return best_ratio
+ def dynamic_preprocess(self, image, min_num=1, max_num=6, image_size=448, use_thumbnail=False):
+ orig_width, orig_height = image.size
+ aspect_ratio = orig_width / orig_height
+ # calculate the existing image aspect ratio
+ target_ratios = set(
+ (i, j) for n in range(min_num, max_num + 1) for i in range(1, n + 1) for j in range(1, n + 1) if
+ i * j <= max_num and i * j >= min_num)
+ target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])
+
+ # find the closest aspect ratio to the target
+ target_aspect_ratio = self.find_closest_aspect_ratio(
+ aspect_ratio, target_ratios, orig_width, orig_height, image_size)
+
+ # calculate the target width and height
+ target_width = image_size * target_aspect_ratio[0]
+ target_height = image_size * target_aspect_ratio[1]
+ blocks = target_aspect_ratio[0] * target_aspect_ratio[1]
+
+ # resize the image
+ resized_img = image.resize((target_width, target_height))
+ processed_images = []
+ for i in range(blocks):
+ box = (
+ (i % (target_width // image_size)) * image_size,
+ (i // (target_width // image_size)) * image_size,
+ ((i % (target_width // image_size)) + 1) * image_size,
+ ((i // (target_width // image_size)) + 1) * image_size
+ )
+ # split the image
+ split_img = resized_img.crop(box)
+ processed_images.append(split_img)
+ assert len(processed_images) == blocks
+ if use_thumbnail and len(processed_images) != 1:
+ thumbnail_img = image.resize((image_size, image_size))
+ processed_images.append(thumbnail_img)
+ return processed_images
+
+ def load_image(self, image_file, input_size=448, max_num=6):
+ if isinstance(image_file, str):
+ image = Image.open(image_file).convert('RGB')
+ else:
+ image = image_file
+ transform = self.build_transform(input_size=input_size)
+ images = self.dynamic_preprocess(image, image_size=input_size, use_thumbnail=True, max_num=max_num)
+ pixel_values = [transform(image) for image in images]
+ pixel_values = torch.stack(pixel_values)
+ return pixel_values
+
+ def get_caption(self, image, prompt, kwargs, use_local, mask):
+ if use_local:
+ image = get_region(image,
+ mask.convert('L'),
+ 255)
+ if image is None:
+ return ""
+ torch.manual_seed(int(time.time()) % 100000)
+ generation_config = dict(
+ num_beams=1,
+ max_new_tokens=4096,
+ do_sample=True
+ )
+ image = self.load_image(image, max_num=6).to(torch.bfloat16).cuda(we.device_id)
+ response = self.model_info['model'].chat(self.model_info['tokenizer'], image, prompt, generation_config=generation_config)
+ response = response.replace("\n", "").strip()
+ return response
diff --git a/scepter/studio/preprocess/processors/image_processors.py b/scepter/studio/preprocess/processors/image_processors.py
index 920a6f8..ce84b1d 100644
--- a/scepter/studio/preprocess/processors/image_processors.py
+++ b/scepter/studio/preprocess/processors/image_processors.py
@@ -1,22 +1,36 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+import numbers
+import numpy as np
+import torch
import torchvision.transforms as TT
+from PIL import Image
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.distribute import we
from scepter.studio.preprocess.processors.base_processor import \
BaseImageProcessor
-__all__ = ['CenterCrop', 'PaddingCrop']
+__all__ = [
+ 'CenterCrop', 'ChangeSample', 'MaskEditSample', 'SwapSample',
+ 'CannyExtractor', 'ColorExtractor', 'InfoDrawContourExtractor',
+ 'DegradationExtractor', 'MidasExtractor', 'DoodleExtractor',
+ 'GrayExtractor', 'InpaintingExtractor', 'OpenposeExtractor',
+ 'OutpaintingExtractor', 'InfoDrawContourExtractor', 'ESAMExtractor',
+ 'InvertExtractor', 'DefaultMaskSample', 'MaskSwapEditSample',
+ 'OutpaintingResize', 'SwapMaskSwapEditSample', 'SourceMaskSample',
+ 'InpaintingSourceExtractor', 'LamaExtractor'
+]
class CenterCrop(BaseImageProcessor):
def __init__(self, cfg, language='en'):
super().__init__(cfg, language=language)
- def __call__(self, image, **kwargs):
+ def process(self, image, height_ratio, width_ratio):
+ if isinstance(image, dict):
+ image = image['background']
w, h = image.size
- height_ratio = kwargs.get('height_ratio', 1)
- width_ratio = kwargs.get('width_ratio', 1)
-
output_h_align_height, output_h_align_width = h, int(h / height_ratio *
width_ratio)
if output_h_align_height * output_h_align_width <= w * h:
@@ -28,6 +42,335 @@ class CenterCrop(BaseImageProcessor):
image = TT.CenterCrop((output_height, output_width))(image)
return image
+ def __call__(self, **kwargs):
+ use_preview = kwargs.get('use_preview', True)
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ target_image = kwargs.get('preview_target_image',
+ None) if use_preview else target_image
+ src_mask = kwargs.get('preview_src_mask',
+ None) if use_preview else src_mask
+ src_image = kwargs.get('preview_src_image',
+ None) if use_preview else src_image
+ height_ratio = kwargs.get('height_ratio', 1)
+ width_ratio = kwargs.get('width_ratio', 1)
+ ret_data = {}
+ if target_image is not None:
+ target_image = self.process(target_image, height_ratio,
+ width_ratio)
+ ret_data['target_image'] = target_image
+ if src_image is not None:
+ src_image = self.process(src_image, height_ratio, width_ratio)
+ ret_data['src_image'] = src_image
+ if src_mask is not None:
+ src_mask = self.process(src_mask, height_ratio, width_ratio)
+ ret_data['src_mask'] = src_mask
+ return ret_data
+
+
+class ChangeSample(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+
+ def __call__(self, **kwargs):
+ # replace file with preview data
+ preview_target_image = kwargs.get('preview_target_image', None)
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption', None)
+ ret_data = {}
+ if preview_target_image is not None:
+ ret_data['target_image'] = preview_target_image['composite']
+ if preview_src_image is not None:
+ ret_data['src_image'] = preview_src_image['background']
+ if preview_src_mask is not None:
+ ret_data['src_mask'] = preview_src_mask['layers'][0]
+ if preview_caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class SourceMaskSample(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+
+ def __call__(self, **kwargs):
+ # replace file with preview data
+ # target_image = kwargs.get('target_image', None)
+ # src_mask = kwargs.get('src_mask', None)
+ # src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get('preview_target_image', None)
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption', None) if use_preview else caption
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['image']
+ prev_src_mask = preview_src_image['mask']
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['image']
+ else:
+ prev_target_image = preview_target_image
+ # process src image and mask according to mask
+ ret_data['src_image'] = None
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ else:
+ prev_src_mask = prev_src_mask.convert('L')
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ # process target image: target_image * mask + src_image * (1 - mask)
+ if prev_target_image is not None:
+ ret_data['target_image'] = prev_target_image
+ else:
+ ret_data['target_image'] = Image.new('RGB', (w, h), (0, 0, 0))
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+class MaskEditSample(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+
+ def __call__(self, **kwargs):
+ # replace file with preview data
+ # target_image = kwargs.get('target_image', None)
+ # src_mask = kwargs.get('src_mask', None)
+ # src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get('preview_target_image', None)
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption',
+ None) if use_preview else caption
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background']
+ prev_src_mask = preview_src_image['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['background']
+ else:
+ prev_target_image = preview_target_image
+ # process src image and mask according to mask
+ ret_data['src_image'] = None
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ else:
+ prev_src_mask = prev_src_mask
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ # process target image: target_image * mask + src_image * (1 - mask)
+ if prev_target_image is not None:
+ prev_target_image = prev_target_image.resize((w, h))
+ ret_data['target_image'] = Image.composite(
+ prev_target_image, prev_src_image, prev_src_mask)
+ else:
+ ret_data['target_image'] = Image.new('RGB', (w, h), (0, 0, 0))
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class SwapMaskSwapEditSample(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+
+ def __call__(self, **kwargs):
+ # replace file with preview data
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get('preview_target_image', None)
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption', None)
+
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background']
+ prev_src_mask = preview_src_image['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['background']
+ else:
+ prev_target_image = preview_target_image
+ # process src image and mask according to mask
+ ret_data['target_image'] = prev_target_image
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ else:
+ prev_src_mask = prev_src_mask
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ prev_src_mask_alter = 255 - np.array(prev_src_mask)
+ prev_src_mask_alter = Image.fromarray(
+ prev_src_mask_alter.astype(np.uint8))
+ # process target image: target_image * mask + src_image * (1 - mask)
+ if prev_target_image is not None:
+ prev_target_image = prev_target_image.resize((w, h))
+ ret_data['src_image'] = Image.composite(
+ prev_target_image, prev_src_image, prev_src_mask_alter)
+ else:
+ ret_data['src_image'] = Image.new('RGB', (w, h), (0, 0, 0))
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class MaskSwapEditSample(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+
+ def __call__(self, **kwargs):
+ # replace file with preview data
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get('preview_target_image', None)
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption', None)
+
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background']
+ prev_src_mask = preview_src_image['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['background']
+ else:
+ prev_target_image = preview_target_image
+ # process src image and mask according to mask
+ ret_data['target_image'] = prev_src_image
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ else:
+ prev_src_mask = prev_src_mask
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ # process target image: target_image * mask + src_image * (1 - mask)
+ if prev_target_image is not None:
+ prev_target_image = prev_target_image.resize((w, h))
+ ret_data['src_image'] = Image.composite(
+ prev_target_image, prev_src_image, prev_src_mask)
+ else:
+ ret_data['src_image'] = Image.new('RGB', (w, h), (0, 0, 0))
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class SwapSample(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+
+ def __call__(self, **kwargs):
+ # replace file with preview data
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ ret_data = {}
+ if target_image is not None and src_image is not None:
+ ret_data['target_image'] = src_image
+ ret_data['src_mask'] = src_mask
+ ret_data['src_image'] = target_image
+ return ret_data
+
+
+class DefaultMaskSample(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+
+ def __call__(self, **kwargs):
+ use_preview = kwargs.get('use_preview', True)
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ preview_target_image = kwargs.get(
+ 'preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.get('preview_src_mask',
+ None) if use_preview else src_mask
+ preview_src_image = kwargs.get('preview_src_image',
+ None) if use_preview else src_image
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background']
+ prev_src_mask = preview_src_image['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ ret_data = {}
+ default_src_mask = np.zeros_like(np.array(prev_src_mask))
+ ret_data['src_image'] = prev_src_image
+ ret_data['src_mask'] = Image.fromarray(
+ default_src_mask.astype(np.uint8))
+ ret_data['target_image'] = preview_target_image
+ return ret_data
+
class PaddingCrop(BaseImageProcessor):
def __init__(self, cfg, language='en'):
@@ -35,3 +378,511 @@ class PaddingCrop(BaseImageProcessor):
def get_caption(self, image, **kwargs):
return image
+
+
+class BaseExtractor(BaseImageProcessor):
+ def __init__(self, cfg, language='en'):
+ super().__init__(cfg, language=language)
+ self.model = cfg.get("MODEL", None)
+ self.model_info = {'device': 'offline', 'model': None}
+
+ def unload_model(self):
+ if self.use_device.lower() == 'gpu':
+ flag, msg = self.unload_model_gpu()
+ else:
+ flag, msg = self.unload_model_cpu()
+ return flag, msg
+
+ def unload_model_cpu(self):
+ super().unload_model()
+ self.model_info['device'] = 'offline'
+ del self.model_info['model']
+ self.model_info['model'] = None
+ return True, ''
+
+ def unload_model_gpu(self):
+ allow_load, msg = self.unload_model_cpu()
+ if self.delete_instance:
+ self.model_info['device'] = 'offline'
+ if self.model_info['model'] is not None:
+ self.model_info['model'] = self.model_info['model'].to('cpu')
+ del self.model_info['model']
+ self.model_info['model'] = None
+ elif (isinstance(self.model_info['device'], numbers.Number)
+ or str(self.model_info['device']).startswith('cuda')):
+ self.model_info['device'] = 'cpu'
+ self.model_info['model'] = self.model_info['model'].to('cpu')
+ torch.cuda.empty_cache()
+ torch.cuda.ipc_collect()
+ return True, ''
+
+ def load_model(self):
+ allow_load, msg = super().load_model()
+ if not allow_load:
+ return allow_load, msg
+ if self.use_device.lower() == 'gpu':
+ flag, msg = self.load_model_gpu()
+ else:
+ flag, msg = self.load_model_cpu()
+ return flag, msg
+
+ def load_model_cpu(self):
+ if self.model_info['device'] == 'offline':
+ self.model_info['device'] = 'cpu'
+ self.model_info['model'] = ANNOTATORS.build(self.model)
+ elif self.model_info['device'] == 'cpu':
+ pass
+ return True, ''
+
+ def load_model_gpu(self):
+ if self.model_info['device'] == 'offline':
+ self.model_info['device'] = 'cpu'
+ self.model_info['model'] = ANNOTATORS.build(self.model).to(
+ we.device_id)
+ elif self.model_info['device'] == 'cpu':
+ self.model_info['device'] = we.device_id
+ self.model_info['model'] = self.model_info['model'].to(
+ we.device_id)
+ return True, ''
+
+ def model_inference(self, model, image, mask=None, **kwargs):
+ image = np.array(image)
+ return model(image)
+
+ def __call__(self, **kwargs):
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get(
+ 'preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.get('preview_src_mask',
+ None) if use_preview else src_mask
+ preview_src_image = kwargs.get('preview_src_image',
+ None) if use_preview else src_image
+ preview_caption = kwargs.get('preview_caption',
+ None) if use_preview else caption
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background'].convert('RGB')
+ prev_src_mask = preview_src_image['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+ # process src image and mask according to mask
+ ret_data['src_image'] = None
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ prev_src_mask = ret_data['src_mask']
+ else:
+ prev_src_mask = prev_src_mask
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ preview_target_image_np = self.model_inference(
+ self.model_info['model'], prev_src_image, prev_src_mask)
+ if len(preview_target_image_np.shape) > 3:
+ preview_target_image = Image.fromarray(
+ preview_target_image_np[0])
+ else:
+ preview_target_image = Image.fromarray(preview_target_image_np)
+ # process target image: target_image * mask + src_image * (1 - mask)
+ if preview_target_image is not None and np.sum(
+ np.array(prev_src_mask)) > 0:
+ preview_target_image = preview_target_image.resize((w, h))
+ ret_data['target_image'] = Image.composite(
+ preview_target_image, prev_src_image, prev_src_mask)
+ else:
+ ret_data['target_image'] = preview_target_image
+ elif preview_target_image is not None:
+ preview_target_image = Image.fromarray(
+ self.model_inference(self.model_info['model'],
+ preview_target_image))
+ ret_data['target_image'] = preview_target_image
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class CannyExtractor(BaseExtractor):
+ pass
+
+
+class ColorExtractor(BaseExtractor):
+ pass
+
+
+class InfoDrawContourExtractor(BaseExtractor):
+ pass
+
+
+class DegradationExtractor(BaseExtractor):
+ pass
+
+
+class MidasExtractor(BaseExtractor):
+ pass
+
+
+class DoodleExtractor(BaseExtractor):
+ pass
+
+
+class GrayExtractor(BaseExtractor):
+ pass
+
+class LamaExtractor(BaseExtractor):
+ def model_inference(self, model, image, mask=None, **kwargs):
+ image = np.array(image)
+ mask = np.array(mask)
+ return model(image, mask)
+
+ def __call__(self, **kwargs):
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get('preview_target_image', None)
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption', None)
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['image']
+ prev_src_mask = preview_src_image['mask']
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+ # process src image and mask according to mask
+ ret_data['src_image'] = prev_src_image
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ prev_src_mask = ret_data['src_mask']
+ else:
+ prev_src_mask = prev_src_mask.convert('L')
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ preview_target_image_np = self.model_inference(
+ self.model_info['model'], prev_src_image, prev_src_mask)
+ if len(preview_target_image_np.shape) > 3:
+ preview_target_image = Image.fromarray(
+ preview_target_image_np[0])
+ else:
+ preview_target_image = Image.fromarray(preview_target_image_np)
+ # process target image: target_image * mask + src_image * (1 - mask)
+ ret_data['target_image'] = preview_target_image
+ elif preview_target_image is not None:
+ preview_target_image = Image.fromarray(
+ self.model_inference(self.model_info['model'],
+ preview_target_image))
+ ret_data['target_image'] = preview_target_image
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+class InpaintingExtractor(BaseExtractor):
+ def model_inference(self, model, image, mask=None, **kwargs):
+ image = np.array(image)
+ mask = np.array(mask) if mask is not None else None
+ return model(image, mask)
+
+ def __call__(self, **kwargs):
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get(
+ 'preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption',
+ None) if use_preview else caption
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background']
+ prev_src_mask = preview_src_image['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ if preview_target_image is not None:
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['background']
+ else:
+ prev_target_image = preview_target_image
+ else:
+ prev_target_image = preview_target_image
+
+ # process src image and mask according to mask
+ ret_data['src_image'] = None
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ prev_src_mask = ret_data['src_mask']
+ else:
+ prev_src_mask = prev_src_mask
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ preview_src_image_np = self.model_inference(
+ self.model_info['model'], prev_src_image, prev_src_mask)
+ if len(preview_src_image_np.shape) > 3:
+ preview_src_image = Image.fromarray(preview_src_image_np[0])
+ else:
+ preview_src_image = Image.fromarray(preview_src_image_np)
+ # process target image: target_image * mask + src_image * (1 - mask)
+ if preview_src_image is not None and np.sum(
+ np.array(prev_src_mask)) > 0:
+ preview_src_image = preview_src_image.resize((w, h))
+ ret_data['src_image'] = Image.composite(
+ preview_src_image, prev_src_image, prev_src_mask)
+ else:
+ ret_data['src_image'] = preview_src_image
+ ret_data['target_image'] = prev_target_image
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class InpaintingSourceExtractor(BaseExtractor):
+ def __call__(self, **kwargs):
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ preview_target_image = kwargs.get('preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.get('preview_src_mask', None) if use_preview else src_mask
+ preview_src_image = kwargs.get('preview_src_image', None) if use_preview else src_image
+ preview_caption = kwargs.get('preview_caption', None) if use_preview else caption
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['image']
+ prev_src_mask = preview_src_mask['image']
+ else:
+ prev_src_image = preview_src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+
+ if preview_target_image is not None:
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['image']
+ else:
+ prev_target_image = preview_target_image
+ else:
+ prev_target_image = preview_target_image
+
+ # process src image and mask according to mask
+ ret_data['src_image'] = None
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ src_mask = ret_data['src_mask']
+ else:
+ src_mask = prev_src_mask.convert('L')
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ ret_data['src_image'] = Image.composite(Image.fromarray(255 - np.array(prev_src_mask)), prev_src_image, src_mask)
+ ret_data['target_image'] = prev_target_image
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class OpenposeExtractor(BaseExtractor):
+ pass
+
+
+class OutpaintingExtractor(BaseExtractor):
+ def model_inference(self,
+ model,
+ image,
+ mask=None,
+ use_mask=False,
+ **kwargs):
+ image = np.array(image)
+ return model(image, mask=mask if use_mask else None, return_mask=True)
+
+ def __call__(self, **kwargs):
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ use_mask = kwargs.get('use_mask', False)
+ # preview_target_image = kwargs.get('preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.get('preview_src_mask', None)
+ preview_src_image = kwargs.get('preview_src_image', None)
+ preview_caption = kwargs.get('preview_caption',
+ None) if use_preview else caption
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image[
+ 'background'] if use_preview else src_image
+ prev_src_mask = preview_src_image['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_image = preview_src_image if use_preview else src_image
+ prev_src_mask = preview_src_mask
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+ # process src image and mask according to mask
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ prev_src_mask = ret_data['src_mask']
+ else:
+ prev_src_mask = prev_src_mask
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ outpaiting_results = self.model_inference(self.model_info['model'],
+ prev_src_image,
+ prev_src_mask,
+ use_mask=use_mask)
+ src_image, target_image, src_mask = (
+ outpaiting_results['src_image'], outpaiting_results['image'],
+ outpaiting_results['mask'])
+ if len(target_image.shape) > 3:
+ preview_target_image = Image.fromarray(target_image[0])
+ else:
+ preview_target_image = Image.fromarray(target_image)
+
+ if len(src_image.shape) > 3:
+ preview_src_image = Image.fromarray(src_image[0])
+ else:
+ preview_src_image = Image.fromarray(src_image)
+
+ if len(src_mask.shape) > 3:
+ preview_src_mask = Image.fromarray(src_mask[0])
+ else:
+ preview_src_mask = Image.fromarray(src_mask)
+
+ ret_data['target_image'] = preview_target_image
+ ret_data['src_image'] = preview_src_image
+ ret_data['src_mask'] = preview_src_mask
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class OutpaintingResize(BaseExtractor):
+ def model_inference(self, model, image, target_image, mask=None, **kwargs):
+ image = np.array(image)
+ return model(image, target_image, mask=mask)
+
+ def __call__(self, **kwargs):
+ target_image = kwargs.get('target_image', None)
+ src_mask = kwargs.get('src_mask', None)
+ src_image = kwargs.get('src_image', None)
+ use_preview = kwargs.get('use_preview', True)
+ caption = kwargs.get('caption', None)
+ use_mask = kwargs.get('use_mask', False)
+ preview_target_image = kwargs.get(
+ 'preview_target_image', None) if use_preview else target_image
+ preview_src_mask = kwargs.get('preview_src_mask',
+ None) if use_preview else src_mask
+ preview_src_image = kwargs.get('preview_src_image',
+ None) if use_preview else src_image
+ preview_caption = kwargs.get('preview_caption',
+ None) if use_preview else caption
+ ret_data = {}
+ if preview_src_image is not None:
+ if isinstance(preview_src_image, dict):
+ prev_src_image = preview_src_image['background']
+ else:
+ prev_src_image = preview_src_image
+ if isinstance(preview_src_mask, dict):
+ prev_src_mask = preview_src_mask['layers'][0].split(
+ )[-1].convert('L')
+ else:
+ prev_src_mask = preview_src_mask
+ if isinstance(preview_target_image, dict):
+ prev_target_image = preview_target_image['background']
+ else:
+ prev_target_image = preview_target_image
+ else:
+ prev_src_image = None
+ prev_src_mask = None
+ prev_target_image = None
+ # process src image and mask according to mask
+ if prev_src_image is not None:
+ w, h = prev_src_image.size
+ if prev_src_mask is None:
+ ret_data['src_mask'] = Image.new('L', (w, h), 0)
+ prev_src_mask = ret_data['src_mask']
+ else:
+ prev_src_mask = prev_src_mask.convert('L')
+ ow, oh = prev_src_mask.size
+ if ow == w and oh == h:
+ ret_data['src_mask'] = prev_src_mask
+ else:
+ ret_data['src_mask'] = Image.new('L', (ow, oh), 0)
+ outpaiting_results = self.model_inference(self.model_info['model'],
+ prev_src_image,
+ prev_target_image,
+ prev_src_mask)
+ src_image = outpaiting_results['src_image']
+ if len(src_image.shape) > 3:
+ preview_src_image = Image.fromarray(src_image[0])
+ else:
+ preview_src_image = Image.fromarray(src_image)
+ ret_data['src_image'] = preview_src_image
+ ret_data['target_image'] = preview_target_image
+ ret_data['src_mask'] = preview_src_mask
+ if caption is not None:
+ ret_data['caption'] = preview_caption
+ return ret_data
+
+
+class InfoDrawAnimeAnnotator(BaseExtractor):
+ pass
+
+
+class ESAMExtractor(BaseExtractor):
+ pass
+
+
+class InvertExtractor(BaseExtractor):
+ pass
diff --git a/scepter/studio/preprocess/utils/data_card.py b/scepter/studio/preprocess/utils/data_card.py
index c2f4b0e..16eb7cc 100644
--- a/scepter/studio/preprocess/utils/data_card.py
+++ b/scepter/studio/preprocess/utils/data_card.py
@@ -242,6 +242,8 @@ class BaseDataCard(object):
illegal_tup.append(f'{range_tup} is illegal, start number '
f'or end number should >= 1. ')
continue
+ start_num = max(0, start_num)
+ end_num = min(len(self), end_num)
if start_num > len(self) or end_num > len(self):
illegal_tup.append(
f'{range_tup} is illegal, start number '
diff --git a/scepter/studio/preprocess/utils/img2img_data_card.py b/scepter/studio/preprocess/utils/img2img_data_card.py
index 1ef3864..74d2785 100644
--- a/scepter/studio/preprocess/utils/img2img_data_card.py
+++ b/scepter/studio/preprocess/utils/img2img_data_card.py
@@ -4,6 +4,7 @@ import csv
import os
import time
+from PIL import Image
from tqdm import tqdm
import gradio as gr
@@ -67,18 +68,50 @@ class Image2ImageDataCard(BaseDataCard):
if 'edit_height' not in cur_data:
self.data[da_idx]['edit_height'] = cur_data['height']
+ if 'src_mask_path' not in cur_data:
+ src_image_name, surfix = os.path.splitext(cur_data['src_relative_path'])
+ src_mask_path = f'{src_image_name}_mask_{int(time.time()*100)}{surfix}'
+ self.meta['local_work_dir'] = self.local_dataset_folder
+ self.meta['work_dir'] = dataset_folder
+ local_src_mask_path = os.path.join(self.meta['local_work_dir'], src_mask_path)
+ src_mask_img = self.default_mask(cur_data['src_height'], cur_data['src_width'])
+ src_mask_img.save(local_src_mask_path)
+ FS.put_object_from_local_file(local_src_mask_path, os.path.join(self.meta['work_dir'], src_mask_path))
+ self.data[da_idx]['src_mask_path'] = os.path.join(self.meta['work_dir'], src_mask_path)
+ self.data[da_idx]['src_mask_relative_path'] = src_mask_path
+ self.data[da_idx]['src_mask_width'] = cur_data['src_width']
+ self.data[da_idx]['src_mask_height'] = cur_data['src_height']
+ else:
+ if self.data[da_idx]['src_mask_path'].startswith(self.meta['work_dir']):
+ self.data[da_idx]['src_mask_relative_path'] = self.data[da_idx]['src_mask_relative_path'].replace(
+ self.meta['work_dir'], ""
+ )
+ if self.data[da_idx]['src_mask_relative_path'].startswith("/"):
+ self.data[da_idx]['src_mask_relative_path'] = self.data[da_idx]['src_mask_relative_path'][1:]
+
+ if 'edit_src_mask_path' not in cur_data:
+ self.data[da_idx]['edit_src_mask_path'] = cur_data[
+ 'src_mask_path']
+ if 'edit_src_mask_relative_path' not in cur_data:
+ self.data[da_idx]['edit_src_mask_relative_path'] = cur_data[
+ 'src_mask_relative_path']
+ if 'edit_src_mask_width' not in cur_data:
+ self.data[da_idx]['edit_src_mask_width'] = cur_data['src_mask_width']
+ if 'edit_src_mask_height' not in cur_data:
+ self.data[da_idx]['edit_src_mask_height'] = cur_data['src_mask_height']
+
if 'edit_src_image_path' not in cur_data:
self.data[da_idx]['edit_src_image_path'] = cur_data[
'src_image_path']
- if 'edit_relative_path' not in cur_data:
+ if 'edit_src_relative_path' not in cur_data:
self.data[da_idx]['edit_src_relative_path'] = cur_data[
'src_relative_path']
- if 'edit_width' not in cur_data:
+ if 'edit_src_width' not in cur_data:
self.data[da_idx]['edit_src_width'] = cur_data['src_width']
- if 'edit_height' not in cur_data:
+ if 'edit_src_height' not in cur_data:
self.data[da_idx]['edit_src_height'] = cur_data[
'src_height']
-
+ self.update_dataset()
def load_from_zip(self, save_file, data_folder, local_dataset_folder):
with FS.get_from(save_file) as local_path:
res = os.popen(
@@ -161,6 +194,13 @@ class Image2ImageDataCard(BaseDataCard):
os.system(f'rm -rf {one_dir}')
return file_list
+ def default_mask(self, h, w):
+ mode = 'L' # 'L' mode is for grayscale images
+ color = 0 # Color value for grayscale (0 = black, 255 = white)
+ # Create the image
+ image = Image.new(mode, (w, h), color)
+ return image
+
def load_train_file(self, file_path, data_folder):
base_folder = os.path.dirname(file_path)
file_list = []
@@ -168,26 +208,49 @@ class Image2ImageDataCard(BaseDataCard):
with open(file_path, 'r') as f:
reader = csv.reader(f)
for row in reader:
- src_image_path, image_path, prompt = row[0], row[1], row[2]
+ if len(row) == 3:
+ src_image_path, image_path, prompt = row[0], row[1], row[2]
+ src_mask_path = None
+ elif len(row) == 4:
+ src_image_path, src_mask_path, image_path, prompt = row[0], row[1], row[2], row[3]
+ else:
+ continue
if image_path == 'Target:FILE':
continue
-
local_src_image_path = os.path.join(base_folder,
src_image_path)
src_w, src_h, src_img = get_image_meta(local_src_image_path)
if src_image_path in image_set:
src_image_name, surfix = os.path.splitext(src_image_path)
- src_image_path = f'{src_image_name}_{int(time.time())}{surfix}'
+ src_image_path = f'{src_image_name}_{int(time.time()*100)}{surfix}'
new_local_image_path = os.path.join(
base_folder, src_image_path)
src_img.save(new_local_image_path)
image_set.add(src_image_path)
+ if src_mask_path is not None and not src_mask_path.strip() == '':
+ local_mask_image_path = os.path.join(base_folder, src_mask_path)
+ src_mask_w, src_mask_h, src_mask_img = get_image_meta(local_mask_image_path)
+ if src_mask_path in image_set:
+ src_mask_name, surfix = os.path.splitext(src_mask_path)
+ src_mask_path = f'{src_mask_name}_{int(time.time()*100)}{surfix}'
+ new_local_src_mask_path = os.path.join(
+ base_folder, src_mask_path)
+ src_img.save(new_local_src_mask_path)
+ image_set.add(src_mask_path)
+ else:
+ src_mask_name, surfix = os.path.splitext(src_image_path)
+ src_mask_path = f'{src_mask_name}_mask_{int(time.time()*100)}{surfix}'
+ local_mask_image_path = os.path.join(base_folder, src_mask_path)
+ src_mask_img = self.default_mask(src_h, src_w)
+ src_mask_img.save(local_mask_image_path)
+ src_mask_w, src_mask_h = src_w, src_h
+
local_image_path = os.path.join(base_folder, image_path)
w, h, img = get_image_meta(local_image_path)
if image_path in image_set:
image_name, surfix = os.path.splitext(image_path)
- image_path = f'{image_name}_{int(time.time())}{surfix}'
+ image_path = f'{image_name}_{int(time.time()*100)}{surfix}'
new_local_image_path = os.path.join(
base_folder, image_path)
img.save(new_local_image_path)
@@ -210,6 +273,14 @@ class Image2ImageDataCard(BaseDataCard):
src_w,
'src_height':
src_h,
+ 'src_mask_path':
+ os.path.join(data_folder, src_mask_path),
+ 'src_mask_relative_path':
+ src_mask_path,
+ 'src_mask_width':
+ src_mask_w,
+ 'src_mask_height':
+ src_mask_h,
'caption':
prompt,
'prefix':
@@ -232,6 +303,14 @@ class Image2ImageDataCard(BaseDataCard):
src_w,
'edit_src_height':
src_h,
+ 'edit_src_mask_path':
+ os.path.join(data_folder, src_mask_path),
+ 'edit_src_mask_relative_path':
+ src_mask_path,
+ 'edit_src_mask_width':
+ src_mask_w,
+ 'edit_src_mask_height':
+ src_h,
})
return file_list
@@ -242,19 +321,20 @@ class Image2ImageDataCard(BaseDataCard):
with FS.get_from(save_file) as local_path:
all_remote_list, all_local_list = [], []
all_src_remote_list, all_src_local_list = [], []
- all_save_list, all_src_save_list = [], []
+ all_src_mask_remote_list, all_src_mask_local_list = [], []
+ all_save_list, all_src_save_list, all_src_mask_save_list = [], [], []
with open(local_path, 'r') as f:
for line in tqdm(f):
line = line.strip()
if line == '':
continue
try:
- src_image_path, src_width, src_height, image_path, width, height, caption = line.split(
- '#;#', 6)
+ src_image_path, src_width, src_height, src_mask_path, src_mask_width, src_mask_height, image_path, width, height, caption = line.split(
+ '#;#', 9)
except Exception:
try:
- src_image_path, src_width, src_height, image_path, width, height, caption = line.split(
- ',', 6)
+ src_image_path, src_width, src_height, src_mask_path, src_mask_width, src_mask_height, image_path, width, height, caption = line.split(
+ ',', 9)
except Exception:
raise gr.Error(
self.components_name.illegal_data_err1)
@@ -275,17 +355,31 @@ class Image2ImageDataCard(BaseDataCard):
self.components_name.illegal_data_err4.format(
src_width, src_height))
+ is_legal, new_src_mask_path, src_mask_prefix = find_prefix(
+ src_mask_path)
+ try:
+ int(src_mask_width), int(src_mask_height)
+ except Exception:
+ raise gr.Error(
+ self.components_name.illegal_data_err4.format(
+ src_mask_width, src_mask_height))
+
+
if not is_legal:
raise gr.Error(
self.components_name.illegal_data_err5.format(
- image_path + ' ' + src_image_path))
+ image_path + ' ' + src_image_path + ' ' + src_mask_path))
relative_path = os.path.join(
'images',
- f'{int(time.time())}_' + image_path.split('/')[-1])
+ f'{int(time.time()*100)}_' + image_path.split('/')[-1])
src_relative_path = os.path.join(
'images',
- f'{int(time.time())}_' + src_image_path.split('/')[-1])
+ f'{int(time.time()*100)}_' + src_image_path.split('/')[-1])
+
+ src_mask_relative_path = os.path.join(
+ 'images',
+ f'{int(time.time()*100)}_' + src_mask_path.split('/')[-1])
all_remote_list.append(new_path)
all_local_list.append(
@@ -299,15 +393,25 @@ class Image2ImageDataCard(BaseDataCard):
all_src_save_list.append(
os.path.join(dataset_folder, src_relative_path))
+ all_src_mask_remote_list.append(new_src_mask_path)
+ all_src_mask_local_list.append(
+ os.path.join(local_dataset_folder, src_mask_relative_path))
+ all_src_mask_save_list.append(
+ os.path.join(dataset_folder, src_mask_relative_path))
+
file_list.append({
'image_path':
os.path.join(dataset_folder, relative_path),
'src_image_path':
os.path.join(dataset_folder, src_relative_path),
+ 'src_mask_path':
+ os.path.join(dataset_folder, src_mask_relative_path),
'relative_path':
relative_path,
'src_relative_path':
src_relative_path,
+ 'src_mask_relative_path':
+ src_mask_relative_path,
'width':
int(width),
'height':
@@ -316,6 +420,10 @@ class Image2ImageDataCard(BaseDataCard):
int(src_width),
'src_height':
int(src_height),
+ 'src_mask_width':
+ int(src_mask_width),
+ 'src_mask_height':
+ int(src_mask_height),
'caption':
caption,
'prefix':
@@ -337,7 +445,15 @@ class Image2ImageDataCard(BaseDataCard):
'edit_src_width':
int(src_width),
'edit_src_height':
- int(src_height)
+ int(src_height),
+ 'edit_src_mask_path':
+ os.path.join(dataset_folder, src_mask_relative_path),
+ 'edit_src_mask_relative_path':
+ src_mask_path,
+ 'edit_src_mask_width':
+ int(src_mask_width),
+ 'edit_src_mask_height':
+ int(src_mask_height)
})
cache_file_list = []
for idx, local_path in enumerate(
@@ -381,13 +497,115 @@ class Image2ImageDataCard(BaseDataCard):
os.remove(local_path)
except Exception:
pass
+ cache_src_mask_file_list = []
+ for idx, local_path in enumerate(
+ FS.get_batch_objects_from(all_src_mask_remote_list)):
+ if local_path is None:
+ raise gr.Error(
+ self.components_name.illegal_data_err6.format(
+ all_src_mask_remote_list[idx]))
+ _ = FS.put_object_from_local_file(local_path,
+ all_src_mask_local_list[idx])
+ cache_src_mask_file_list.append(local_path)
+
+ for local_path, target_path, flg in FS.put_batch_objects_to(
+ cache_src_mask_file_list, all_src_mask_save_list):
+ if not flg:
+ raise gr.Error(
+ self.components_name.illegal_data_err7.format(local_path))
+ if os.path.exists(local_path):
+ try:
+ os.remove(local_path)
+ except Exception:
+ pass
return file_list
+ def apply_changes(self):
+ edit_index_list = self.edit_list
+ for index in edit_index_list:
+ one_data = self.data[index]
+
+ relative_image_path = one_data['relative_path']
+ local_image_path = os.path.join(self.meta['local_work_dir'],
+ relative_image_path)
+ image_path = one_data['image_path']
+
+ edit_relative_image_path = one_data.get('edit_relative_path',
+ one_data['relative_path'])
+ local_edit_image_path = os.path.join(self.meta['local_work_dir'],
+ edit_relative_image_path)
+ if not relative_image_path == edit_relative_image_path:
+ try:
+ os.rename(local_edit_image_path, local_image_path)
+ except Exception as e:
+ msg = f'Apply edited image failed, error is {e}'
+ return False, msg
+ FS.put_object_from_local_file(local_image_path, image_path)
+ try:
+ os.remove(local_edit_image_path)
+ except Exception:
+ pass
+
+ src_relative_path = one_data['src_relative_path']
+ local_src_path = os.path.join(self.meta['local_work_dir'],
+ src_relative_path)
+ src_image_path = one_data['src_image_path']
+
+ edit_src_relative_path = one_data.get('edit_src_relative_path',
+ one_data['src_relative_path'])
+ local_edit_src_image_path = os.path.join(self.meta['local_work_dir'],
+ edit_src_relative_path)
+ if not src_relative_path == edit_src_relative_path:
+ try:
+ os.rename(local_edit_src_image_path, local_src_path)
+ except Exception as e:
+ msg = f'Apply edited image failed, error is {e}'
+ return False, msg
+ FS.put_object_from_local_file(local_src_path, src_image_path)
+ try:
+ os.remove(local_edit_image_path)
+ except Exception:
+ pass
+
+ src_mask_relative_path = one_data['src_mask_relative_path']
+ local_src_mask_path = os.path.join(self.meta['local_work_dir'],
+ src_mask_relative_path)
+ src_mask_path = one_data['src_mask_path']
+
+ edit_src_mask_relative_path = one_data.get('edit_src_mask_relative_path',
+ one_data['src_mask_relative_path'])
+ local_edit_src_mask_path = os.path.join(self.meta['local_work_dir'],
+ edit_src_mask_relative_path)
+ if not src_mask_relative_path == edit_src_mask_relative_path:
+ try:
+ os.rename(local_edit_src_mask_path, local_src_mask_path)
+ except Exception as e:
+ msg = f'Apply edited image failed, error is {e}'
+ return False, msg
+ FS.put_object_from_local_file(local_src_mask_path, src_mask_path)
+ try:
+ os.remove(local_edit_src_mask_path)
+ except Exception:
+ pass
+
+
+ self.data[index]['edit_relative_path'] = one_data['relative_path']
+ self.data[index]['edit_image_path'] = one_data['image_path']
+ self.data[index]['edit_src_relative_path'] = one_data['src_relative_path']
+ self.data[index]['edit_src_image_path'] = one_data['src_image_path']
+ self.data[index]['edit_src_mask_relative_path'] = one_data['src_mask_relative_path']
+ self.data[index]['edit_src_mask_path'] = one_data['src_mask_path']
+ self.data[index]['caption'] = one_data['edit_caption']
+ self.data[index]['width'] = one_data['edit_width']
+ self.data[index]['height'] = one_data['edit_height']
+ self.update_dataset()
+ return True, ''
+
def write_train_file(self):
file_list = self.meta['file_list']
with open(self.local_train_file, 'w') as f:
writer = csv.writer(f)
- writer.writerow(['Source:FILE', 'Target:FILE', 'Prompt'])
+ writer.writerow(['Source:FILE', 'SourceMASK:FILE', 'Target:FILE', 'Prompt'])
for one_file in file_list:
relative_file = one_file['relative_path']
if relative_file.startswith('/'):
@@ -397,8 +615,12 @@ class Image2ImageDataCard(BaseDataCard):
if src_relative_file.startswith('/'):
src_relative_file = src_relative_file[1:]
+ src_mask_relative_file = one_file['src_mask_relative_path']
+ if src_mask_relative_file.startswith('/'):
+ src_mask_relative_file = src_mask_relative_file[1:]
+
writer.writerow(
- [src_relative_file, relative_file, one_file['caption']])
+ [src_relative_file, src_mask_relative_file, relative_file, one_file['caption'].strip().replace("\n", "")])
FS.put_object_from_local_file(self.local_train_file, self.train_file)
def write_data_file(self):
@@ -409,9 +631,13 @@ class Image2ImageDataCard(BaseDataCard):
prefix=one_file['prefix'])
is_flag, src_file_path = del_prefix(one_file['src_image_path'],
prefix=one_file['prefix'])
- f.write('{}#;#{}#;#{}#;#{}#;#{}#;#{}#;#{}\n'.format(
+ is_flag, src_mask_file_path = del_prefix(one_file['src_mask_path'],
+ prefix=one_file['prefix'])
+ f.write('{}#;#{}#;#{}#;#{}#;#{}#;#{}#;#{}#;#{}#;#{}#;#{}\n'.format(
src_file_path, one_file['src_width'],
- one_file['src_height'], file_path, one_file['width'],
+ one_file['src_height'], src_mask_file_path,
+ one_file['src_mask_width'], one_file['src_mask_height'],
+ file_path, one_file['width'],
one_file['height'], one_file['caption']))
FS.put_object_from_local_file(self.local_save_file_list,
self.save_file_list)
@@ -422,10 +648,9 @@ class Image2ImageDataCard(BaseDataCard):
save_folder = os.path.join(local_work_dir, 'images')
os.makedirs(save_folder, exist_ok=True)
-
w, h = image.size
relative_path = os.path.join(
- 'images', f'{imagehash.phash(image)}_{int(time.time())}.jpg')
+ 'images', f'{imagehash.phash(image)}_{int(time.time()*100)}.jpg')
image_path = os.path.join(work_dir, relative_path)
local_image_path = os.path.join(local_work_dir, relative_path)
image.save(local_image_path)
@@ -434,12 +659,26 @@ class Image2ImageDataCard(BaseDataCard):
src_image = kwargs.pop('src_image')
src_w, src_h = src_image.size
src_relative_path = os.path.join(
- 'images', f'{imagehash.phash(src_image)}_{int(time.time())}.jpg')
+ 'images', f'{imagehash.phash(src_image)}_{int(time.time()*100)}.jpg')
src_image_path = os.path.join(work_dir, src_relative_path)
local_src_image_path = os.path.join(local_work_dir, src_relative_path)
src_image.save(local_src_image_path)
FS.put_object_from_local_file(local_src_image_path, src_image_path)
+ if 'src_mask' not in kwargs:
+ src_mask_image = None
+ else:
+ src_mask_image = kwargs.pop('src_mask')
+ if src_mask_image is None:
+ src_mask_image = self.default_mask(src_h, src_w)
+ src_mask_w, src_mask_h = src_mask_image.size
+ src_mask_relative_path = os.path.join(
+ 'images', f'{imagehash.phash(src_mask_image)}_{int(time.time()*100)}.jpg')
+ src_mask_path = os.path.join(work_dir, src_mask_relative_path)
+ local_src_mask_path = os.path.join(local_work_dir, src_mask_relative_path)
+ src_mask_image.save(local_src_mask_path)
+ FS.put_object_from_local_file(local_src_mask_path, src_mask_path)
+
self.data.append({
'image_path': image_path,
'relative_path': relative_path,
@@ -449,6 +688,10 @@ class Image2ImageDataCard(BaseDataCard):
'src_relative_path': src_relative_path,
'src_width': src_w,
'src_height': src_h,
+ 'src_mask_path': src_mask_path,
+ 'src_mask_relative_path': src_mask_relative_path,
+ 'src_mask_width': src_mask_w,
+ 'src_mask_height': src_mask_h,
'caption': caption,
'prefix': '',
'edit_caption': caption,
@@ -459,9 +702,12 @@ class Image2ImageDataCard(BaseDataCard):
'edit_src_image_path': src_image_path,
'edit_src_relative_path': src_relative_path,
'edit_src_width': src_w,
- 'edit_src_height': src_h
+ 'edit_src_height': src_h,
+ 'edit_src_mask_path': src_mask_path,
+ 'edit_src_mask_relative_path': src_mask_relative_path,
+ 'edit_src_mask_width': src_mask_w,
+ 'edit_src_mask_height': src_mask_h
})
-
self.set_cursor(len(self.meta['file_list']) - 1)
self.update_dataset()
return True
@@ -485,6 +731,13 @@ class Image2ImageDataCard(BaseDataCard):
except Exception:
print(f'remove file {local_src_file} error')
+ local_src_mask_file = os.path.join(self.meta['local_work_dir'],
+ current_file['src_mask_relative_path'])
+ try:
+ os.remove(local_src_mask_file)
+ except Exception:
+ print(f'remove file {local_src_mask_file} error')
+
if self.cursor >= len(self.meta['file_list']):
self.set_cursor(0)
if self.cursor < 0:
diff --git a/scepter/studio/preprocess/utils/txt2img_data_card.py b/scepter/studio/preprocess/utils/txt2img_data_card.py
index 2e45313..df68ffe 100644
--- a/scepter/studio/preprocess/utils/txt2img_data_card.py
+++ b/scepter/studio/preprocess/utils/txt2img_data_card.py
@@ -334,7 +334,7 @@ class Text2ImageDataCard(BaseDataCard):
relative_file = one_file['relative_path']
if relative_file.startswith('/'):
relative_file = relative_file[1:]
- writer.writerow([relative_file, one_file['caption']])
+ writer.writerow([relative_file, one_file['caption'].strip().replace("\n", "")])
FS.put_object_from_local_file(self.local_train_file, self.train_file)
def write_data_file(self):
@@ -346,7 +346,7 @@ class Text2ImageDataCard(BaseDataCard):
f.write('{}#;#{}#;#{}#;#{}\n'.format(file_path,
one_file['width'],
one_file['height'],
- one_file['caption']))
+ one_file['caption'].strip().replace("\n", "")))
FS.put_object_from_local_file(self.local_save_file_list,
self.save_file_list)
diff --git a/scepter/studio/self_train/scripts/trainer.py b/scepter/studio/self_train/scripts/trainer.py
index 7680254..4ca7924 100644
--- a/scepter/studio/self_train/scripts/trainer.py
+++ b/scepter/studio/self_train/scripts/trainer.py
@@ -319,6 +319,17 @@ class TrainManager():
self.task_queue.remove(task_name)
else:
pass
+
+ status_file = os.path.join(self.work_dir, task_name, 'status.json')
+ with FS.get_from(status_file) as local_status:
+ task_status = json.load(open(local_status, 'r'))
+ task_status['status'] = 'failed'
+ task_status['end_time'] = time.time()
+ task_status['msg'] = 'Task has been killed.'
+ with FS.put_to(status_file) as local_path:
+ json.dump(task_status,
+ open(local_path, 'w'),
+ ensure_ascii=False)
# modify the task's status met
def __del__(self):
diff --git a/scepter/studio/self_train/self_train_ui/model_ui.py b/scepter/studio/self_train/self_train_ui/model_ui.py
index 871de5b..7cac932 100644
--- a/scepter/studio/self_train/self_train_ui/model_ui.py
+++ b/scepter/studio/self_train/self_train_ui/model_ui.py
@@ -6,9 +6,8 @@ import os
import queue
import time
-import yaml
-
import gradio as gr
+import yaml
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.studio.self_train.self_train_ui.component_names import ModelUIName
@@ -17,6 +16,7 @@ from scepter.studio.utils.uibase import UIBase
refresh_symbol = '\U0001f504' # 🔄
delete_symbol = '\U0001f5d1' # 🗑️
+stop_symbol = '\U000025A0'
add_symbol = '\U00002795' # ➕
confirm_symbol = '\U00002714' # ✔️
@@ -31,32 +31,41 @@ class ModelUI(UIBase):
self.default_model_name = 'step-last'
self.old_default_model_name = 'checkpoint.pth'
self.delete_folder_queue = queue.Queue()
- self.model_list = []
- # self.model_list.extend(self.base_model_info.get('model_choices', []))
- have_model_list = []
+ self.user_level_model_list = {}
if not self.work_dir.endswith('/'):
self.work_dir += '/'
+ self.load_history()
+ self.is_debug = is_debug
+ self.component_names = ModelUIName(language)
+
+ def load_history(self, user_name='admin'):
+ have_model_list = []
for one_dir in FS.walk_dir(self.work_dir):
if one_dir.startswith(self.work_dir):
one_dir = one_dir[len(self.work_dir):]
if not os.path.isdir(os.path.join(self.work_dir, one_dir)):
continue
- if len(one_dir.split('/')) > 1:
- continue
- if ('@' in one_dir
- and (os.path.exists(
- os.path.join(self.work_dir, one_dir, 'checkpoints',
- self.default_model_name))
- or os.path.exists(
- os.path.join(self.work_dir, one_dir,
- self.old_default_model_name)))):
- if len(one_dir.split('@')) > 4:
- have_model_list.append([one_dir, one_dir.split('@')[-1]])
- have_model_list.sort(key=lambda x: -int(x[-1][:-3]))
- self.model_list.extend([v[0].split('/')[-1]
- for v in have_model_list][:50])
- self.is_debug = is_debug
- self.component_names = ModelUIName(language)
+ params_file = os.path.join(self.work_dir, one_dir, 'params.json')
+ if FS.exists(params_file):
+ try:
+ params = json.load(open(params_file, 'r'))
+ except:
+ continue
+ model_user_name = params.get('user_name', 'example')
+ if model_user_name == user_name or model_user_name in [
+ 'example'
+ ] or user_name == 'admin':
+ # parse params
+ status_file = os.path.join(self.work_dir, one_dir,
+ 'status.json')
+ if FS.exists(status_file):
+ status = json.load(open(status_file, 'r'))
+ status['model_name'] = one_dir
+ have_model_list.append(status)
+ have_model_list.sort(key=lambda x: x['start_time'])
+ self.user_level_model_list[user_name] = [
+ v['model_name'] for v in have_model_list
+ ][:100]
def get_ckpt_list(self, output_model):
all_ckpt_list = []
@@ -95,7 +104,7 @@ class ModelUI(UIBase):
with gr.Column(scale=7, min_width=0):
self.output_model_name = gr.Dropdown(
label=self.component_names.output_model_name,
- choices=self.model_list,
+ choices=self.user_level_model_list.get('admin', []),
value=None,
show_label=False,
container=False,
@@ -112,6 +121,8 @@ class ModelUI(UIBase):
self.refresh_model_gbtn = gr.Button(value=refresh_symbol)
with gr.Column(scale=1, min_width=0):
self.add_model_gbtn = gr.Button(value=add_symbol)
+ with gr.Column(scale=1, min_width=0):
+ self.stop_model_gbtn = gr.Button(value=stop_symbol)
with gr.Column(scale=1, min_width=0):
self.delete_model_gbtn = gr.Button(value=delete_symbol)
@@ -221,23 +232,26 @@ class ModelUI(UIBase):
outputs=[self.extra_model_panel],
queue=False)
- def confirm_add(model_name):
+ def confirm_add(model_name, login_user_name):
model_folder = os.path.join(self.work_dir, model_name)
have_model = os.path.exists(model_folder)
if not have_model:
gr.Error(self.component_names.model_err5.format(model_name))
- if model_name not in self.model_list and have_model:
- self.model_list.append(model_name)
- return gr.Row(visible=False), gr.Dropdown(choices=self.model_list,
- value=model_name)
+ self.load_history(login_user_name)
+ if model_name not in self.user_level_model_list[
+ login_user_name] and have_model:
+ self.user_level_model_list[login_user_name].append(model_name)
+ return gr.Row(visible=False), gr.Dropdown(
+ choices=self.user_level_model_list[login_user_name],
+ value=model_name)
self.confirm_add.click(
fn=confirm_add,
- inputs=[self.extra_model_txt],
+ inputs=[self.extra_model_txt, manager.user_name],
outputs=[self.extra_model_panel, self.output_model_name],
queue=False)
- def refresh_model(model_name):
+ def refresh_model(model_name, login_user_name):
if not self.delete_folder_queue.empty():
try:
del_folder = self.delete_folder_queue.get_nowait()
@@ -252,29 +266,33 @@ class ModelUI(UIBase):
ckpt_list = self.get_ckpt_list(model_name)
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
ret_gallery = ckpt_name_change(model_name, ckpt_value)
+ self.load_history(login_user_name)
return (message, gr.Column(visible=status in ('running',
'success')),
- gr.Dropdown(choices=self.model_list, value=model_name),
+ gr.Dropdown(choices=self.user_level_model_list.get(
+ login_user_name, []),
+ value=model_name),
gr.Dropdown(choices=ckpt_list,
value=ckpt_value), ret_gallery)
- self.refresh_model_gbtn.click(fn=refresh_model,
- inputs=[self.output_model_name],
- outputs=[
- self.log_message,
- self.export_log_panel,
- self.output_model_name,
- self.output_ckpt_name,
- self.eval_gallery
- ],
- queue=False)
+ self.refresh_model_gbtn.click(
+ fn=refresh_model,
+ inputs=[self.output_model_name, manager.user_name],
+ outputs=[
+ self.log_message, self.export_log_panel,
+ self.output_model_name, self.output_ckpt_name,
+ self.eval_gallery
+ ],
+ queue=False)
- def delete_model(model_name):
+ def delete_model(model_name, login_user_name):
index = 0
trainer_ui.trainer_ins.stop_task(model_name)
- if model_name in self.model_list:
- index = self.model_list.index(model_name)
- self.model_list.remove(model_name)
+ self.load_history(login_user_name)
+ model_list = self.user_level_model_list.get(login_user_name, [])
+ if model_name in model_list:
+ index = model_list.index(model_name)
+ model_list.remove(model_name)
folder = os.path.join(self.work_dir, model_name)
for _ in range(1):
if os.path.exists(folder):
@@ -285,20 +303,56 @@ class ModelUI(UIBase):
else:
break
self.delete_folder_queue.put_nowait(folder)
- if index <= len(self.model_list) - 1:
- model_name = self.model_list[index]
- elif len(self.model_list) > 0:
- model_name = self.model_list[0]
+ if index <= len(model_list) - 1:
+ model_name = model_list[index]
+ elif len(model_list) > 0:
+ model_name = model_list[-1]
else:
model_name = None
+ self.user_level_model_list[login_user_name] = model_list
+ return gr.Dropdown(choices=model_list, value=model_name)
- return gr.Dropdown(choices=self.model_list, value=model_name)
+ self.delete_model_gbtn.click(
+ fn=delete_model,
+ inputs=[self.output_model_name, manager.user_name],
+ outputs=[self.output_model_name])
- self.delete_model_gbtn.click(fn=delete_model,
- inputs=[self.output_model_name],
- outputs=[self.output_model_name])
+ def stop_model(model_name, login_user_name):
+ trainer_ui.trainer_ins.stop_task(model_name)
+ if not self.delete_folder_queue.empty():
+ try:
+ del_folder = self.delete_folder_queue.get_nowait()
+ if os.path.exists(del_folder):
+ os.system(f'rm -rf {del_folder}')
+ self.delete_folder_queue.put_nowait(del_folder)
+ except Exception:
+ pass
+ message = trainer_ui.trainer_ins.get_log(model_name)
+ status = trainer_ui.trainer_ins.get_status(model_name)
+ ckpt_list = self.get_ckpt_list(model_name)
+ ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
+ ret_gallery = ckpt_name_change(model_name, ckpt_value)
+ self.load_history(login_user_name)
+ return (message, gr.Column(visible=status in ('running',
+ 'success')),
+ gr.Dropdown(choices=self.user_level_model_list.get(
+ login_user_name, []),
+ value=model_name),
+ gr.Dropdown(choices=ckpt_list,
+ value=ckpt_value), ret_gallery)
+
+ self.stop_model_gbtn.click(
+ fn=stop_model,
+ inputs=[self.output_model_name, manager.user_name],
+ outputs=[
+ self.log_message, self.export_log_panel,
+ self.output_model_name, self.output_ckpt_name,
+ self.eval_gallery
+ ])
def go_to_inferece(output_model, output_ckpt_name):
+ if output_ckpt_name is None or output_ckpt_name.strip() == '':
+ raise gr.Error(message="The model file doesn't exist.")
params_path = os.path.join(self.work_dir, output_model,
'params.json')
if os.path.exists(params_path):
@@ -492,3 +546,17 @@ class ModelUI(UIBase):
inputs=[self.output_model_name],
outputs=[self.export_url],
queue=False)
+
+ def model_name_change(login_user_name):
+ self.load_history(login_user_name)
+ model_list = self.user_level_model_list.get(login_user_name, [])
+ if len(model_list) > 0:
+ model_name = model_list[-1]
+ else:
+ model_name = ''
+ return gr.Dropdown(choices=model_list, value=model_name)
+
+ manager.user_name.change(model_name_change,
+ inputs=[manager.user_name],
+ outputs=[self.output_model_name],
+ queue=False)
diff --git a/scepter/studio/self_train/self_train_ui/trainer_ui.py b/scepter/studio/self_train/self_train_ui/trainer_ui.py
index b2ff280..6889fe2 100644
--- a/scepter/studio/self_train/self_train_ui/trainer_ui.py
+++ b/scepter/studio/self_train/self_train_ui/trainer_ui.py
@@ -200,23 +200,23 @@ class TrainerUI(UIBase):
self.lora_alpha = gr.Number(
label='LoRA Alpha',
value=self.para_data.get(
- 'lora_alpha', 256),
+ 'lora_alpha', 4),
interactive=True)
self.lora_rank = gr.Number(
label='LoRA Rank',
- value=self.para_data.get('lora_rank', 256),
+ value=self.para_data.get('lora_rank', 4),
interactive=True)
with gr.Row(visible=text_lora_visible
) as self.text_lora_param:
self.text_lora_alpha = gr.Number(
label='Text LoRA Alpha',
value=self.para_data.get(
- 'text_lora_alpha', 256),
+ 'text_lora_alpha', 4),
interactive=True)
self.text_lora_rank = gr.Number(
label='Text LoRA Rank',
value=self.para_data.get(
- 'text_lora_rank', 256),
+ 'text_lora_rank', 4),
interactive=True)
self.sce_ratio = gr.Slider(
label='SCE Ratio',
@@ -613,7 +613,7 @@ class TrainerUI(UIBase):
lora_alpha, lora_rank, text_lora_alpha, text_lora_rank,
sce_ratio, enable_resolution_bucket,
min_bucket_resolution, max_bucket_resolution,
- bucket_resolution_steps, bucket_no_upscale):
+ bucket_resolution_steps, bucket_no_upscale, user_name):
# Check Cuda
if not torch.cuda.is_available() and not self.is_debug:
raise gr.Error(self.component_names.training_err1)
@@ -621,6 +621,7 @@ class TrainerUI(UIBase):
if work_name == 'custom' or work_name is None or work_name == '':
raise gr.Error(self.component_names.training_err4)
work_dir = os.path.join(self.work_dir_pre, work_name)
+ login_user_name = user_name
self.current_train_model = work_name
if os.path.exists(work_dir) or os.path.exists(
f'.flag/{work_name}.tmp'):
@@ -895,11 +896,14 @@ class TrainerUI(UIBase):
_ = prepare_train_config()
self.trainer_ins.start_task(work_name)
message += self.trainer_ins.get_log(work_name)
- if work_name not in inference_ui.model_list:
- inference_ui.model_list.append(work_name)
+ if user_name not in inference_ui.user_level_model_list:
+ inference_ui.user_level_model_list[user_name] = []
+ if work_name not in inference_ui.user_level_model_list[user_name]:
+ inference_ui.user_level_model_list[user_name].append(work_name)
gr.Info('Start Training!' + message)
- return gr.Dropdown(choices=inference_ui.model_list,
- value=work_name)
+ return gr.Dropdown(
+ choices=inference_ui.user_level_model_list[user_name],
+ value=work_name)
self.training_button.click(
run_train,
@@ -916,7 +920,7 @@ class TrainerUI(UIBase):
self.text_lora_alpha, self.text_lora_rank, self.sce_ratio,
self.enable_resolution_bucket, self.min_bucket_resolution,
self.max_bucket_resolution, self.bucket_resolution_steps,
- self.bucket_no_upscale
+ self.bucket_no_upscale, manager.user_name
],
outputs=[inference_ui.output_model_name],
queue=True)
diff --git a/scepter/studio/self_train/utils/config_parser.py b/scepter/studio/self_train/utils/config_parser.py
index 6198876..d9b0e3a 100644
--- a/scepter/studio/self_train/utils/config_parser.py
+++ b/scepter/studio/self_train/utils/config_parser.py
@@ -21,7 +21,7 @@ def build_meta_index(meta_cfg, config_file):
for key in paras_keys:
if key not in para:
print(
- f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
+ f'Para key {key} not defined in {config_file} META/RARAS[{idx}]' # noqa
)
assert key in para
tuner_type[para['TUNER']] = para
@@ -47,7 +47,7 @@ def build_meta_index_control(meta_cfg, config_file):
for key in control_paras_keys:
if key not in para:
print(
- f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
+ f'Para key {key} not defined in {config_file} META/RARAS[{idx}]' # noqa
)
assert key in para
control_type[para['CONTROL_MODE']] = para
@@ -137,11 +137,10 @@ def get_all_config(config_root, global_meta):
def get_default(config_dict):
ret_data = {}
- # 默认的模型
ret_data['model_choices'] = config_dict['choices']
ret_data['model_default'] = config_dict['default']
default_version_cfg = config_dict.get(config_dict['default'], None)
- # 默认的版本
+
if default_version_cfg is None:
return ret_data
ret_data['version_choices'] = default_version_cfg['choices']
@@ -288,11 +287,10 @@ def get_control_para_by_model_version(config_dict, model_name, version):
def get_control_default(config_dict):
ret_data = {}
- # 默认的模型
ret_data['model_choices'] = config_dict['choices']
ret_data['model_default'] = config_dict['default']
default_version_cfg = config_dict.get(config_dict['default'], None)
- # 默认的版本
+
if default_version_cfg is None:
return ret_data
ret_data['version_choices'] = default_version_cfg['choices']
diff --git a/scepter/studio/tuner_manager/manager_ui/browser_ui.py b/scepter/studio/tuner_manager/manager_ui/browser_ui.py
index c88cdf4..2562384 100644
--- a/scepter/studio/tuner_manager/manager_ui/browser_ui.py
+++ b/scepter/studio/tuner_manager/manager_ui/browser_ui.py
@@ -20,7 +20,6 @@ from scepter.studio.tuner_manager.utils.path import (
is_valid_modelscope_filename)
from scepter.studio.tuner_manager.utils.yaml import save_yaml
from scepter.studio.utils.uibase import UIBase
-from swift import push_to_hub
def wget_file(file_url, save_file):
@@ -63,29 +62,37 @@ class BrowserUI(UIBase):
for tuner in self.saved_tuners:
first_level = f"{tuner['BASE_MODEL']}-{tuner['TUNER_TYPE']}"
second_level = f"{tuner['NAME']}"
+ login_user_name = tuner.get('USER_NAME', 'admin')
+ if login_user_name not in self.saved_tuners_category:
+ self.saved_tuners_category[login_user_name] = OrderedDict()
local_tuner_work_dir, _ = FS.map_to_local(tuner['MODEL_PATH'])
if not os.path.exists(local_tuner_work_dir):
FS.get_dir_to_local_dir(tuner['MODEL_PATH'])
- update_2level_dict(self.saved_tuners_category,
+ update_2level_dict(self.saved_tuners_category[login_user_name],
{first_level: {
second_level: tuner
}})
def category_to_saved_tuners(self):
self.saved_tuners = []
- for k, v in self.saved_tuners_category.items():
- for kk, vv in v.items():
- self.saved_tuners.append(vv)
+ for login_user_name, user_model in self.saved_tuners_category.items():
+ for k, v in user_model.items():
+ for kk, vv in v.items():
+ self.saved_tuners.append(vv)
- def get_choices_and_values(self):
- diffusion_models_choice = list(self.saved_tuners_category.keys())
+ def get_choices_and_values(self, login_user_name='admin'):
+ diffusion_models_choice = list(
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).keys())
diffusion_model = diffusion_models_choice[0] if len(
diffusion_models_choice) > 0 else None
tuner_models_choice = []
tuner_model = None
if diffusion_model:
tuner_models_choice = list(
- self.saved_tuners_category.get(diffusion_model, {}).keys())
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).get(
+ diffusion_model, {}).keys())
tuner_model = tuner_models_choice[0] if len(
tuner_models_choice) > 0 else None
return diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model
@@ -380,6 +387,7 @@ class BrowserUI(UIBase):
params = Config.get_plain_cfg(meta.get('PARAMS', {}))
params['work_dir'] = ''
params['work_name'] = ''
+ params['USER_NAME'] = ''
save_yaml(
{
'PARAMS':
@@ -432,8 +440,8 @@ class BrowserUI(UIBase):
return tar_path, tuner_example_path, enable_share
def save_tuner(self, manager, tuner_name, new_name, tuner_desc,
- tuner_example, tuner_prompt_example, base_model,
- tuner_type):
+ tuner_example, tuner_prompt_example, base_model, tuner_type,
+ login_user_name):
is_legal, msg = self.check_new_name(new_name)
if not is_legal:
raise gr.Error('Save failed because ' + msg)
@@ -450,8 +458,10 @@ class BrowserUI(UIBase):
'@'.join(tuner_name.split('@')[:-1]),
'checkpoints', steps)
else:
- source_dir = self.saved_tuners_category.get(sub_dir, {}).get(
- tuner_name, {}).get('MODEL_PATH', '')
+ source_dir = self.saved_tuners_category.get(
+ login_user_name,
+ OrderedDict()).get(sub_dir, {}).get(tuner_name,
+ {}).get('MODEL_PATH', '')
if source_dir == '':
raise gr.Error(self.component_names.model_err4 + tuner_name)
@@ -465,6 +475,7 @@ class BrowserUI(UIBase):
'SOURCE': 'self_train',
'DESCRIPTION': tuner_desc,
'BASE_MODEL': base_model,
+ 'USER_NAME': login_user_name,
'MODEL_PATH': model_dir,
'TUNER_TYPE': tuner_type,
'PROMPT_EXAMPLE': tuner_prompt_example,
@@ -479,19 +490,25 @@ class BrowserUI(UIBase):
pipeline_ins = pipeline_level_modules[new_tuner['BASE_MODEL']]
now_diffusion_model = f"{new_tuner['BASE_MODEL']}_{pipeline_ins.diffusion_model['name']}"
- custom_tuner_choices = self.add_tuner(new_tuner, manager,
- now_diffusion_model)
+ custom_tuner_choices = self.add_tuner(new_tuner,
+ manager,
+ now_diffusion_model,
+ login_user_name=login_user_name)
gr.Info('Successfully save tuner model!')
- return (gr.update(choices=list(self.saved_tuners_category.keys()),
+ return (gr.update(choices=list(
+ self.saved_tuners_category.get(login_user_name, {}).keys()),
value=sub_dir),
gr.update(choices=list(
- self.saved_tuners_category.get(sub_dir, {}).keys()),
+ self.saved_tuners_category.get(login_user_name,
+ {}).get(sub_dir,
+ {}).keys()),
value=new_name), gr.Text(value=new_name),
gr.Dropdown(choices=custom_tuner_choices))
- def add_tuner(self, new_tuner, manager, now_diffusion_model):
+ def add_tuner(self, new_tuner, manager, now_diffusion_model,
+ login_user_name):
self.saved_tuners.append(new_tuner)
self.saved_tuners_to_category()
with FS.put_to(self.yaml) as local_path:
@@ -529,9 +546,11 @@ class BrowserUI(UIBase):
return custom_tunner_choices
def delete_tuner(self, first_level, second_level, manager,
- now_diffusion_model):
- self.saved_tuners_category, del_tuner = delete_2level_dict(
- self.saved_tuners_category, first_level, second_level)
+ now_diffusion_model, login_user_name):
+ self.saved_tuners_category[
+ login_user_name], del_tuner = delete_2level_dict(
+ self.saved_tuners_category.get(login_user_name, OrderedDict()),
+ first_level, second_level)
self.category_to_saved_tuners()
save_yaml({'TUNERS': self.saved_tuners}, self.yaml)
@@ -557,16 +576,24 @@ class BrowserUI(UIBase):
return custom_tuner_choices
- def update_tuner_info(self, base_model, tuner_name, update_items):
+ def update_tuner_info(self, base_model, tuner_name, update_items,
+ login_user_name):
+ if login_user_name not in self.saved_tuners_category:
+ self.saved_tuners_category[login_user_name] = OrderedDict()
+ self.saved_tuners_category[login_user_name][
+ base_model] = OrderedDict()
+ if tuner_name not in self.saved_tuners_category[login_user_name][
+ base_model]:
+ raise gr.Error(f'{tuner_name} not found in {base_model}')
self.saved_tuners_category[base_model][tuner_name].update(update_items)
self.category_to_saved_tuners()
with FS.put_to(self.yaml) as local_path:
save_yaml({'TUNERS': self.saved_tuners}, local_path)
def set_callbacks(self, manager, info_ui):
- def refresh_browser():
+ def refresh_browser(login_user_name):
diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model = self.get_choices_and_values(
- )
+ login_user_name=login_user_name)
return (gr.Dropdown(choices=diffusion_models_choice,
value=diffusion_model),
gr.Dropdown(choices=tuner_models_choice,
@@ -574,28 +601,31 @@ class BrowserUI(UIBase):
self.refresh_button.click(
refresh_browser,
- inputs=[],
+ inputs=[manager.user_name],
outputs=[self.diffusion_models, self.tuner_models])
- def diffusion_model_change(diffusion_model):
+ def diffusion_model_change(diffusion_model, login_user_name):
choices = list(
- self.saved_tuners_category.get(diffusion_model, {}).keys())
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).get(
+ diffusion_model, {}).keys())
return gr.Dropdown(choices=choices,
value=choices[-1] if len(choices) > 0 else None)
- self.diffusion_models.change(diffusion_model_change,
- inputs=[self.diffusion_models],
- outputs=[self.tuner_models],
- queue=True)
+ self.diffusion_models.change(
+ diffusion_model_change,
+ inputs=[self.diffusion_models, manager.user_name],
+ outputs=[self.tuner_models],
+ queue=True)
- def tuner_model_change(tuner_model, diffusion_model):
+ def tuner_model_change(tuner_model, diffusion_model, login_user_name):
if tuner_model is None:
# fix refresh bug
return (gr.Text(), gr.Text(), gr.Text(), gr.Text(), gr.Text(),
gr.Image(), gr.Text(), gr.Text(), gr.Text())
- tuner_info = self.saved_tuners_category[diffusion_model][
- tuner_model]
+ tuner_info = self.saved_tuners_category.get(
+ login_user_name, OrderedDict())[diffusion_model][tuner_model]
local_model_dir, _ = FS.map_to_local(tuner_info['MODEL_PATH'])
image_path = tuner_info.get('IMAGE_PATH', None)
if image_path is not None:
@@ -619,34 +649,37 @@ class BrowserUI(UIBase):
gr.Text(value=tuner_info.get('HUGGINGFACE_URL', ''),
interactive=False))
- self.tuner_models.change(
- tuner_model_change,
- inputs=[self.tuner_models, self.diffusion_models],
- outputs=[
- info_ui.tuner_name,
- info_ui.new_name,
- info_ui.tuner_type,
- info_ui.base_model,
- info_ui.tuner_desc,
- info_ui.tuner_example,
- info_ui.tuner_prompt_example,
- info_ui.ms_url,
- info_ui.hf_url,
- ],
- queue=False)
+ self.tuner_models.change(tuner_model_change,
+ inputs=[
+ self.tuner_models, self.diffusion_models,
+ manager.user_name
+ ],
+ outputs=[
+ info_ui.tuner_name,
+ info_ui.new_name,
+ info_ui.tuner_type,
+ info_ui.base_model,
+ info_ui.tuner_desc,
+ info_ui.tuner_example,
+ info_ui.tuner_prompt_example,
+ info_ui.ms_url,
+ info_ui.hf_url,
+ ],
+ queue=False)
def save_tuner_func(tuner_name, new_name, tuner_desc, tuner_example,
- tuner_prompt_example, base_model, tuner_type):
+ tuner_prompt_example, base_model, tuner_type,
+ login_user_name):
return self.save_tuner(manager, tuner_name, new_name, tuner_desc,
tuner_example, tuner_prompt_example,
- base_model, tuner_type)
+ base_model, tuner_type, login_user_name)
self.save_button.click(
save_tuner_func,
inputs=[
info_ui.tuner_name, info_ui.new_name, info_ui.tuner_desc,
info_ui.tuner_example, info_ui.tuner_prompt_example,
- info_ui.base_model, info_ui.tuner_type
+ info_ui.base_model, info_ui.tuner_type, manager.user_name
],
outputs=[
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
@@ -655,21 +688,24 @@ class BrowserUI(UIBase):
queue=True)
def delete_tuner(tuner_name, tuner_type, base_model,
- now_diffusion_model):
+ now_diffusion_model, login_user_name):
first_level = f'{base_model}-{tuner_type}'
second_level = f'{tuner_name}'
custom_tuner_choices = self.delete_tuner(first_level, second_level,
manager,
- now_diffusion_model)
- return (gr.Dropdown(
- choices=list(self.saved_tuners_category.keys()),
- value=None), gr.Dropdown(choices=custom_tuner_choices))
+ now_diffusion_model,
+ login_user_name)
+ return (gr.Dropdown(choices=list(
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).keys()),
+ value=None),
+ gr.Dropdown(choices=custom_tuner_choices))
self.delete_button.click(
delete_tuner,
inputs=[
info_ui.tuner_name, info_ui.tuner_type, info_ui.base_model,
- manager.inference.model_manage_ui.diffusion_model
+ manager.inference.model_manage_ui.diffusion_model, manager.user_name
],
outputs=[
self.diffusion_models,
@@ -771,20 +807,23 @@ class BrowserUI(UIBase):
queue=False)
def push_to_modelscope(ms_sdk, username, private, base_model_name,
- tuner_model_name):
+ tuner_model_name, login_user_name):
+ from swift import push_to_hub
gr.Info('Start uploading tuner model to ModelScope!')
- if (isinstance(base_model_name, list) and len(base_model_name) == 0
- ) or (isinstance(tuner_model_name, list)
- and len(tuner_model_name)
- == 0) or tuner_model_name is None or (
- base_model_name not in self.saved_tuners_category
- and tuner_model_name
- not in self.saved_tuners_category[base_model_name]):
+ if (isinstance(base_model_name, list)
+ and len(base_model_name) == 0) or (
+ isinstance(tuner_model_name, list)
+ and len(tuner_model_name) == 0
+ ) or tuner_model_name is None or (
+ base_model_name not in self.saved_tuners_category.get(
+ login_user_name, {}) and tuner_model_name
+ not in self.saved_tuners_category[login_user_name]
+ [base_model_name]):
raise gr.Error(
'Please save model first or select a valid base model name.'
)
- tuner = self.saved_tuners_category[base_model_name][
- tuner_model_name]
+ tuner = self.saved_tuners_category[login_user_name][
+ base_model_name][tuner_model_name]
enable_share = tuner.get('ENABLE_SHARE', True)
if enable_share:
@@ -802,7 +841,7 @@ class BrowserUI(UIBase):
with open(local_readme, 'r') as f:
rc = f.read()
rc = rc.replace(r'{MODEL_URL}', ms_url)
- rc = rc.replace(r'{USER_NAME}', username)
+ rc = rc.replace(r'{login_user_name}', username)
with open(local_readme, 'w') as f:
f.write(rc)
local_configuration = os.path.join(local_ckpt_path,
@@ -840,12 +879,12 @@ class BrowserUI(UIBase):
fn=push_to_modelscope,
inputs=[
self.ms_sdk, self.ms_export_username, self.ms_model_private,
- self.diffusion_models, self.tuner_models
+ self.diffusion_models, self.tuner_models, manager.user_name
],
outputs=[info_ui.ms_url, self.export_setting],
queue=True)
- def pull_from_modelscope(modelid, username):
+ def pull_from_modelscope(modelid, username, login_user_name):
gr.Info('Start pulling tuner model from ModelScope to Local!')
src_path = f'ms://{username}/{modelid}'
local_work_dir, _ = FS.map_to_local(src_path)
@@ -880,6 +919,7 @@ class BrowserUI(UIBase):
new_tuner = {
'NAME': new_name,
'NAME_ZH': new_name,
+ 'USER_NAME': login_user_name,
'SOURCE': 'modelscope',
'DESCRIPTION': meta.get('DESCRIPTION', ''),
'BASE_MODEL': base_model,
@@ -895,8 +935,11 @@ class BrowserUI(UIBase):
pipeline_ins = pipeline_level_modules[new_tuner['BASE_MODEL']]
now_diffusion_model = f"{new_tuner['BASE_MODEL']}_{pipeline_ins.diffusion_model['name']}"
- custom_tuner_choices = self.add_tuner(new_tuner, manager,
- now_diffusion_model)
+ custom_tuner_choices = self.add_tuner(
+ new_tuner,
+ manager,
+ now_diffusion_model,
+ login_user_name=login_user_name)
update_items = {
'MODELSCOPE_URL':
@@ -906,11 +949,15 @@ class BrowserUI(UIBase):
new_name,
update_items=update_items)
- return (gr.update(choices=list(self.saved_tuners_category.keys()),
+ return (gr.update(choices=list(
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).keys()),
value=tuner_category),
gr.update(choices=list(
- self.saved_tuners_category.get(tuner_category,
- {}).keys()),
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).get(
+ tuner_category,
+ {}).keys()),
value=new_name), gr.Text(value=new_name),
gr.Dropdown(choices=custom_tuner_choices),
gr.update(visible=False))
@@ -920,6 +967,7 @@ class BrowserUI(UIBase):
inputs=[
self.ms_modelid,
self.ms_import_username,
+ manager.user_name
],
outputs=[
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
@@ -929,20 +977,22 @@ class BrowserUI(UIBase):
queue=True)
def push_to_huggingface(sdk, username, private, base_model_name,
- tuner_model_name):
+ tuner_model_name, login_user_name):
gr.Info('Start uploading tuner model to HuggingFace!')
- if (isinstance(base_model_name, list) and len(base_model_name) == 0
- ) or (isinstance(tuner_model_name, list)
- and len(tuner_model_name)
- == 0) or tuner_model_name is None or (
- base_model_name not in self.saved_tuners_category
- and tuner_model_name
- not in self.saved_tuners_category[base_model_name]):
+ if (
+ isinstance(base_model_name, list)
+ and len(base_model_name) == 0
+ ) or (isinstance(tuner_model_name, list) and len(tuner_model_name)
+ == 0) or tuner_model_name is None or (
+ base_model_name not in self.saved_tuners_category.get(
+ login_user_name, OrderedDict())
+ and tuner_model_name not in self.
+ saved_tuners_category[login_user_name][base_model_name]):
raise gr.Error(
'Please save model first or select a valid base model name.'
)
- tuner = self.saved_tuners_category[base_model_name][
- tuner_model_name]
+ tuner = self.saved_tuners_category[login_user_name][
+ base_model_name][tuner_model_name]
enable_share = tuner.get('ENABLE_SHARE', True)
if enable_share:
@@ -960,7 +1010,7 @@ class BrowserUI(UIBase):
with open(local_readme, 'r') as f:
rc = f.read()
rc = rc.replace(r'{MODEL_URL}', hf_url)
- rc = rc.replace(r'{USER_NAME}', username)
+ rc = rc.replace(r'{login_user_name}', username)
with open(local_readme, 'w') as f:
f.write(rc)
local_configuration = os.path.join(ckpt_path,
@@ -999,12 +1049,12 @@ class BrowserUI(UIBase):
fn=push_to_huggingface,
inputs=[
self.hf_sdk, self.hf_export_username, self.hf_model_private,
- self.diffusion_models, self.tuner_models
+ self.diffusion_models, self.tuner_models, manager.user_name
],
outputs=[info_ui.hf_url, self.export_setting],
queue=True)
- def pull_from_huggingface(modelid, username, hf_sdk):
+ def pull_from_huggingface(modelid, username, hf_sdk, login_user_name):
gr.Info('Start pulling tuner model from HuggingFace to Local!')
src_path = f'{username}/{modelid}'
@@ -1058,7 +1108,8 @@ class BrowserUI(UIBase):
now_diffusion_model = f"{new_tuner['BASE_MODEL']}_{pipeline_ins.diffusion_model['name']}"
custom_tuner_choices = self.add_tuner(new_tuner, manager,
- now_diffusion_model)
+ now_diffusion_model,
+ login_user_name)
update_items = {
'HUGGINGFACE_URL':
@@ -1066,13 +1117,18 @@ class BrowserUI(UIBase):
}
self.update_tuner_info(tuner_category,
new_name,
- update_items=update_items)
+ update_items=update_items,
+ login_user_name=login_user_name)
- return (gr.update(choices=list(self.saved_tuners_category.keys()),
+ return (gr.update(choices=list(
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).keys()),
value=tuner_category),
gr.update(choices=list(
- self.saved_tuners_category.get(tuner_category,
- {}).keys()),
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).get(
+ tuner_category,
+ {}).keys()),
value=new_name), gr.Text(value=new_name),
gr.Dropdown(choices=custom_tuner_choices),
gr.update(visible=False))
@@ -1080,9 +1136,8 @@ class BrowserUI(UIBase):
self.hf_import_submit.click(
fn=pull_from_huggingface,
inputs=[
- self.hf_modelid,
- self.hf_import_username,
- self.hf_sdk2,
+ self.hf_modelid, self.hf_import_username, self.hf_sdk2,
+ manager.user_name
],
outputs=[
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
@@ -1129,8 +1184,8 @@ class BrowserUI(UIBase):
FS.get_dir_to_local_dir(model_dir, local_model_dir)
return model_dir, local_model_dir
- def upload_zip(file_path, file_url, tuner_name, base_model,
- tuner_type):
+ def upload_zip(file_path, file_url, tuner_name, base_model, tuner_type,
+ login_user_name):
sub_dir = f'{base_model}-{tuner_type}'
sub_work_dir = os.path.join(self.work_dir, sub_dir)
if not FS.exists(sub_work_dir):
@@ -1220,6 +1275,7 @@ class BrowserUI(UIBase):
'NAME': tuner_name,
'NAME_ZH': tuner_name,
'SOURCE': 'self_train',
+ 'USER_NAME': login_user_name,
'DESCRIPTION': tuner_desc,
'BASE_MODEL': base_model,
'MODEL_PATH': model_dir,
@@ -1239,15 +1295,22 @@ class BrowserUI(UIBase):
pipeline_ins = pipeline_level_modules[new_tuner['BASE_MODEL']]
now_diffusion_model = f"{new_tuner['BASE_MODEL']}_{pipeline_ins.diffusion_model['name']}"
- custom_tuner_choices = self.add_tuner(new_tuner, manager,
- now_diffusion_model)
+ custom_tuner_choices = self.add_tuner(
+ new_tuner,
+ manager,
+ now_diffusion_model,
+ login_user_name=login_user_name)
gr.Info(self.component_names.upload_success)
return (gr.Dropdown(choices=list(
- self.saved_tuners_category.keys()),
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).keys()),
value=sub_dir),
gr.Dropdown(choices=list(
- self.saved_tuners_category.get(sub_dir, {}).keys()),
+ self.saved_tuners_category.get(login_user_name,
+ OrderedDict()).get(
+ sub_dir,
+ {}).keys()),
value=tuner_name), gr.Text(value=tuner_name),
gr.Text(value=tuner_type), gr.Text(value=base_model),
gr.Text(value=tuner_desc),
@@ -1260,7 +1323,7 @@ class BrowserUI(UIBase):
upload_zip,
inputs=[
self.file_path, self.file_url, self.upload_tuner_name,
- self.upload_base_models, self.upload_tuner_type
+ self.upload_base_models, self.upload_tuner_type, manager.user_name
],
outputs=[
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
diff --git a/scepter/studio/tuner_manager/manager_ui/info_ui.py b/scepter/studio/tuner_manager/manager_ui/info_ui.py
index 0e92315..8fc4900 100644
--- a/scepter/studio/tuner_manager/manager_ui/info_ui.py
+++ b/scepter/studio/tuner_manager/manager_ui/info_ui.py
@@ -2,11 +2,11 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os
-
-import yaml
+from collections import OrderedDict
import gradio as gr
import torch
+import yaml
from safetensors.torch import save_file
from scepter.modules.utils.config import Config
from scepter.modules.utils.directory import get_md5
@@ -146,8 +146,9 @@ class InfoUI(UIBase):
def set_callbacks(self, manager, browser_ui):
# def set_callbacks(self, manager):
def go_to_inferece(new_name, tuner_desc, tuner_prompt_example,
- tuner_type, base_model):
- all_tuners = manager.tuner_manager.browser_ui.saved_tuners_category
+ tuner_type, base_model, login_user_name):
+ all_tuners = manager.tuner_manager.browser_ui.saved_tuners_category.get(
+ login_user_name, OrderedDict())
sub_dir = f'{base_model}-{tuner_type}'
current_model = all_tuners.get(sub_dir, {}).get(new_name, None)
if current_model is None:
@@ -241,7 +242,7 @@ class InfoUI(UIBase):
go_to_inferece,
inputs=[
self.new_name, self.tuner_desc, self.tuner_prompt_example,
- self.tuner_type, self.base_model
+ self.tuner_type, self.base_model, manager.user_name
],
outputs=[
manager.tabs, manager.inference.infer_info,
@@ -252,8 +253,9 @@ class InfoUI(UIBase):
],
queue=True)
- def export_zip(tuner_name, base_model, tuner_type):
- all_tuners = manager.tuner_manager.browser_ui.saved_tuners_category
+ def export_zip(tuner_name, base_model, tuner_type, login_user_name):
+ all_tuners = manager.tuner_manager.browser_ui.saved_tuners_category.get(
+ login_user_name, OrderedDict())
sub_dir = f'{base_model}-{tuner_type}'
current_model = all_tuners.get(sub_dir, {}).get(tuner_name, None)
if current_model is None:
@@ -276,9 +278,11 @@ class InfoUI(UIBase):
gr.Info(self.component_names.save_end)
return gr.File(value=local_zip, visible=True)
- def export_safetensors(tuner_name, base_model, tuner_type):
+ def export_safetensors(tuner_name, base_model, tuner_type,
+ login_user_name):
# only support sd1.5
- all_tuners = manager.tuner_manager.browser_ui.saved_tuners_category
+ all_tuners = manager.tuner_manager.browser_ui.saved_tuners_category.get(
+ login_user_name, OrderedDict())
sub_dir = f'{base_model}-{tuner_type}'
current_model = all_tuners.get(sub_dir, {}).get(tuner_name, None)
if current_model is None:
@@ -325,12 +329,15 @@ class InfoUI(UIBase):
gr.Info(self.component_names.save_end)
return gr.File(value=local_zip, visible=True)
- def export_file(tuner_name, base_model, tuner_type, download_type):
+ def export_file(tuner_name, base_model, tuner_type, download_type,
+ login_user_name):
gr.Info(self.component_names.save_start)
if download_type == '.zip':
- return export_zip(tuner_name, base_model, tuner_type)
+ return export_zip(tuner_name, base_model, tuner_type,
+ login_user_name)
elif download_type == '.safetensors':
- return export_safetensors(tuner_name, base_model, tuner_type)
+ return export_safetensors(tuner_name, base_model, tuner_type,
+ login_user_name)
def change_visible():
return gr.update(visible=True)
@@ -343,24 +350,26 @@ class InfoUI(UIBase):
self.download_confirm.click(export_file,
inputs=[
self.tuner_name, self.base_model,
- self.tuner_type, self.download_select
+ self.tuner_type, self.download_select,
+ manager.user_name
],
outputs=[self.export_url],
queue=True)
def save_tuner_func(tuner_name, new_name, tuner_desc, tuner_example,
- tuner_prompt_example, base_model, tuner_type):
+ tuner_prompt_example, base_model, tuner_type,
+ login_user_name):
return browser_ui.save_tuner(manager, tuner_name, new_name,
tuner_desc, tuner_example,
tuner_prompt_example, base_model,
- tuner_type)
+ tuner_type, login_user_name)
self.save_bt2.click(save_tuner_func,
inputs=[
self.tuner_name, self.new_name,
self.tuner_desc, self.tuner_example,
self.tuner_prompt_example, self.base_model,
- self.tuner_type
+ self.tuner_type, manager.user_name
],
outputs=[
browser_ui.diffusion_models,
diff --git a/scepter/tools/local.sh b/scepter/tools/local.sh
deleted file mode 100644
index 6a0ef5b..0000000
--- a/scepter/tools/local.sh
+++ /dev/null
@@ -1,2 +0,0 @@
-export CUDA_VISIBLE_DEVICES=0
-python -W ignore scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_xl_1024_align.yaml
diff --git a/scepter/tools/run_inference.py b/scepter/tools/run_inference.py
index 349da04..685172a 100644
--- a/scepter/tools/run_inference.py
+++ b/scepter/tools/run_inference.py
@@ -7,7 +7,6 @@ import sys
import numpy as np
import torch
-import torch.cuda.amp as amp
import torchvision.transforms as TT
from PIL import Image
from scepter.modules.solver.registry import SOLVERS
@@ -26,6 +25,7 @@ if os.path.exists('__init__.py'):
def run_task(cfg):
+ import torch.cuda.amp as amp
std_logger = get_logger(name='scepter')
solver = SOLVERS.build(cfg.SOLVER, logger=std_logger)
solver.set_up()
diff --git a/scepter/tools/webui.py b/scepter/tools/webui.py
index b055240..a413d6a 100644
--- a/scepter/tools/webui.py
+++ b/scepter/tools/webui.py
@@ -94,6 +94,7 @@ if __name__ == '__main__':
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
+ print('init home page success!')
if ifid == 'preprocess':
from scepter.studio.preprocess.preprocess import PreprocessUI
@@ -101,6 +102,7 @@ if __name__ == '__main__':
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
+ print('init preprocess success!')
if ifid == 'self_train':
from scepter.studio.self_train.self_train import SelfTrainUI
@@ -108,12 +110,14 @@ if __name__ == '__main__':
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
+ print('init self-train success!')
if ifid == 'tuner_manager':
from scepter.studio.tuner_manager.tuner_manager import TunerManagerUI
interface = TunerManagerUI(info['CONFIG'],
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
+ print('init tuner-manager success!')
if ifid == 'inference':
from scepter.studio.inference.inference import InferenceUI
@@ -121,6 +125,7 @@ if __name__ == '__main__':
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
+ print('init inference success!')
if ifid == '':
pass # TODO: Add New Features
if interface:
diff --git a/scepter/version.py b/scepter/version.py
index 5205e57..84f60fe 100644
--- a/scepter/version.py
+++ b/scepter/version.py
@@ -1,7 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
-__version__ = '1.0.3'
+__version__ = '1.1.0'
version_info = tuple(int(x) for x in __version__.split('.')[0:3])
diff --git a/scepter/workflow/__init__.py b/scepter/workflow/__init__.py
new file mode 100644
index 0000000..01bdf34
--- /dev/null
+++ b/scepter/workflow/__init__.py
@@ -0,0 +1,22 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+from .model_node import ModelNode
+from .note_node import NoteNode
+from .parameter_node import ParameterNode
+from .mantras_node import MantrasNode
+from .tuner_node import TunerNode
+from .control_node import ControlNode
+
+NODE_MAPPINGS = {
+ 'ModelNode': ('🪄 ScepterModel~', ModelNode),
+ 'NoteNode': ('🪄 ScepterNote~', NoteNode),
+ 'ParameterNode': ('🪄 ScepterParameter~', ParameterNode),
+ 'MantrasNode': ('🪄 ScepterMantra~', MantrasNode),
+ 'TunerNode': ('🪄 ScepterTuner~', TunerNode),
+ 'ControlNode': ('🪄 ScepterControl~', ControlNode)
+}
+
+NODE_CLASS_MAPPINGS = {k : v[1] for k, v in NODE_MAPPINGS.items()}
+NODE_DISPLAY_NAME_MAPPINGS = {k : v[0] for k, v in NODE_MAPPINGS.items()}
+
+__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
diff --git a/scepter/workflow/config/annotator.yaml b/scepter/workflow/config/annotator.yaml
new file mode 100644
index 0000000..658dbbb
--- /dev/null
+++ b/scepter/workflow/config/annotator.yaml
@@ -0,0 +1,29 @@
+ANNOTATORS:
+ -
+ NAME: "CannyAnnotator"
+ TYPE: Canny
+ IS_DEFAULT: True
+ -
+ NAME: "HedAnnotator"
+ PRETRAINED_MODEL: "ms://iic/scepter_scedit@annotator/ckpts/ControlNetHED.pth"
+ TYPE: Hed
+ IS_DEFAULT: False
+ -
+ NAME: "OpenposeAnnotator"
+ BODY_MODEL_PATH: "ms://iic/scepter_scedit@annotator/ckpts/body_pose_model.pth"
+ HAND_MODEL_PATH: "ms://iic/scepter_scedit@annotator/ckpts/hand_pose_model.pth"
+ TYPE: Openpose
+ IS_DEFAULT: False
+ -
+ NAME: "MidasDetector"
+ PRETRAINED_MODEL: "ms://iic/scepter_scedit@annotator/ckpts/dpt_hybrid-midas-501f0c75.pth"
+ TYPE: Midas
+ IS_DEFAULT: False
+ -
+ NAME: "ColorAnnotator"
+ TYPE: Color
+ IS_DEFAULT: False
+ -
+ NAME: "InvertAnnotator"
+ TYPE: Invert-Preprocess
+ IS_DEFAULT: False
diff --git a/scepter/workflow/config/control_model.yaml b/scepter/workflow/config/control_model.yaml
new file mode 100644
index 0000000..3839594
--- /dev/null
+++ b/scepter/workflow/config/control_model.yaml
@@ -0,0 +1,32 @@
+CONTROLLERS:
+ # SD_XL1.0
+ - NAME: SD_XL1.0_canny
+ NAME_ZH:
+ DESCRIPTION:
+ BASE_MODEL: SD_XL1.0
+ TYPE: Canny
+ MODEL_PATH: ms://iic/scepter_scedit@controllable_model/SD_XL1.0/canny_control
+ - NAME: SD_XL1.0_color
+ NAME_ZH:
+ DESCRIPTION:
+ BASE_MODEL: SD_XL1.0
+ TYPE: Color
+ MODEL_PATH: ms://iic/scepter_scedit@controllable_model/SD_XL1.0/color_control
+ - NAME: SD_XL1.0_depth
+ NAME_ZH:
+ DESCRIPTION:
+ BASE_MODEL: SD_XL1.0
+ TYPE: Midas
+ MODEL_PATH: ms://iic/scepter_scedit@controllable_model/SD_XL1.0/depth_control
+ - NAME: SD_XL1.0_hed
+ NAME_ZH:
+ DESCRIPTION:
+ BASE_MODEL: SD_XL1.0
+ TYPE: Hed
+ MODEL_PATH: ms://iic/scepter_scedit@controllable_model/SD_XL1.0/hed_control
+ - NAME: SD_XL1.0_openpose
+ NAME_ZH:
+ DESCRIPTION:
+ BASE_MODEL: SD_XL1.0
+ TYPE: Openpose
+ MODEL_PATH: ms://iic/scepter_scedit@controllable_model/SD_XL1.0/pose_control
diff --git a/scepter/workflow/config/flux1.0_dev_pro.yaml b/scepter/workflow/config/flux1.0_dev_pro.yaml
new file mode 100644
index 0000000..3e1f5d9
--- /dev/null
+++ b/scepter/workflow/config/flux1.0_dev_pro.yaml
@@ -0,0 +1,366 @@
+NAME: FLUX1.0_DEV
+IS_DEFAULT: False
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[1024, 1024]]
+ INPUT:
+ IMAGE:
+ ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
+ TARGET_SIZE_AS_TUPLE: [1024, 1024]
+ PROMPT: ""
+ NEGATIVE_PROMPT:
+ DEFAULT: ""
+ VISIBLE: False
+ PROMPT_PREFIX: ""
+ SAMPLE:
+ VALUES: ["flow_eluer"]
+ DEFAULT: "flow_eluer"
+ SAMPLE_STEPS: 50
+ GUIDE_SCALE: 3.5
+ GUIDE_RESCALE:
+ DEFAULT: 0.0
+ VISIBLE: False
+ DISCRETIZATION:
+ VALUES: []
+ DEFAULT:
+ VISIBLE: False
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: bfloat16
+ INPUT: ["LATENT"]
+ PARAS:
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: bfloat16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["PROMPT"]
+#
+MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ NOISE_SCHEDULER:
+ NAME: FlowMatchSigmaScheduler
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ LOGIT_MEAN: 0.0
+ LOGIT_STD: 1.0
+ MODE_SCALE: 1.29
+ #
+ DIFFUSION_MODEL:
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
+ IN_CHANNELS: 64
+ HIDDEN_SIZE: 3072
+ NUM_HEADS: 24
+ AXES_DIM: [ 16, 56, 56 ]
+ THETA: 10000
+ VEC_IN_DIM: 768
+ GUIDANCE_EMBED: True
+ CONTEXT_IN_DIM: 4096
+ MLP_RATIO: 4.0
+ QKV_BIAS: True
+ DEPTH: 19
+ DEPTH_SINGLE_BLOCKS: 38
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5PlusClipFluxEmbedder
+ T5_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: T5EncoderModel
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
+ HF_TOKENIZER_CLS: T5Tokenizer
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
+ MAX_LENGTH: 512
+ OUTPUT_KEY: last_hidden_state
+ D_TYPE: bfloat16
+ BATCH_INFER: False
+ CLEAN: whitespace
+ #
+ CLIP_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: CLIPTextModel
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
+ MAX_LENGTH: 77
+ OUTPUT_KEY: pooler_output
+ D_TYPE: bfloat16
+ BATCH_INFER: True
+ CLEAN: whitespace
+#
+MODEL_LOCAL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ NOISE_SCHEDULER:
+ NAME: FlowMatchSigmaScheduler
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ LOGIT_MEAN: 0.0
+ LOGIT_STD: 1.0
+ MODE_SCALE: 1.29
+ #
+ DIFFUSION_MODEL:
+ NAME: Flux
+ PRETRAINED_MODEL: models/scepter/FLUX.1-dev/flux1-dev.safetensors
+ IN_CHANNELS: 64
+ HIDDEN_SIZE: 3072
+ NUM_HEADS: 24
+ AXES_DIM: [ 16, 56, 56 ]
+ THETA: 10000
+ VEC_IN_DIM: 768
+ GUIDANCE_EMBED: True
+ CONTEXT_IN_DIM: 4096
+ MLP_RATIO: 4.0
+ QKV_BIAS: True
+ DEPTH: 19
+ DEPTH_SINGLE_BLOCKS: 38
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: models/scepter/FLUX.1-dev/ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5PlusClipFluxEmbedder
+ T5_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: T5EncoderModel
+ MODEL_PATH: models/scepter/FLUX.1-dev/text_encoder_2/
+ HF_TOKENIZER_CLS: T5Tokenizer
+ TOKENIZER_PATH: models/scepter/FLUX.1-dev/tokenizer_2/
+ MAX_LENGTH: 512
+ OUTPUT_KEY: last_hidden_state
+ D_TYPE: bfloat16
+ BATCH_INFER: False
+ CLEAN: whitespace
+ #
+ CLIP_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: CLIPTextModel
+ MODEL_PATH: models/scepter/FLUX.1-dev/text_encoder/
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ TOKENIZER_PATH: models/scepter/FLUX.1-dev/tokenizer/
+ MAX_LENGTH: 77
+ OUTPUT_KEY: pooler_output
+ D_TYPE: bfloat16
+ BATCH_INFER: True
+ CLEAN: whitespace
+#
+MODEL_HF:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ NOISE_SCHEDULER:
+ NAME: FlowMatchSigmaScheduler
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ LOGIT_MEAN: 0.0
+ LOGIT_STD: 1.0
+ MODE_SCALE: 1.29
+ #
+ DIFFUSION_MODEL:
+ NAME: Flux
+ PRETRAINED_MODEL: hf://black-forest-labs/FLUX.1-dev@flux1-dev.safetensors
+ IN_CHANNELS: 64
+ HIDDEN_SIZE: 3072
+ NUM_HEADS: 24
+ AXES_DIM: [ 16, 56, 56 ]
+ THETA: 10000
+ VEC_IN_DIM: 768
+ GUIDANCE_EMBED: True
+ CONTEXT_IN_DIM: 4096
+ MLP_RATIO: 4.0
+ QKV_BIAS: True
+ DEPTH: 19
+ DEPTH_SINGLE_BLOCKS: 38
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: hf://black-forest-labs/FLUX.1-dev@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5PlusClipFluxEmbedder
+ T5_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: T5EncoderModel
+ MODEL_PATH: hf://black-forest-labs/FLUX.1-dev@text_encoder_2/
+ HF_TOKENIZER_CLS: T5Tokenizer
+ TOKENIZER_PATH: hf://black-forest-labs/FLUX.1-dev@tokenizer_2/
+ MAX_LENGTH: 512
+ OUTPUT_KEY: last_hidden_state
+ D_TYPE: bfloat16
+ BATCH_INFER: False
+ CLEAN: whitespace
+ #
+ CLIP_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: CLIPTextModel
+ MODEL_PATH: hf://black-forest-labs/FLUX.1-dev@text_encoder/
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ TOKENIZER_PATH: hf://black-forest-labs/FLUX.1-dev@tokenizer/
+ MAX_LENGTH: 77
+ OUTPUT_KEY: pooler_output
+ D_TYPE: bfloat16
+ BATCH_INFER: True
+ CLEAN: whitespace
\ No newline at end of file
diff --git a/scepter/workflow/config/flux1.0_schnell_pro.yaml b/scepter/workflow/config/flux1.0_schnell_pro.yaml
new file mode 100644
index 0000000..b86a958
--- /dev/null
+++ b/scepter/workflow/config/flux1.0_schnell_pro.yaml
@@ -0,0 +1,384 @@
+NAME: FLUX1.0_SCHNELL
+IS_DEFAULT: False
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[1024, 1024]]
+ INPUT:
+ IMAGE:
+ ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
+ TARGET_SIZE_AS_TUPLE: [1024, 1024]
+ PROMPT: ""
+ NEGATIVE_PROMPT:
+ DEFAULT: ""
+ VISIBLE: False
+ PROMPT_PREFIX: ""
+ SAMPLE:
+ VALUES: ["flow_eluer"]
+ DEFAULT: "flow_eluer"
+ SAMPLE_STEPS: 4
+ GUIDE_SCALE: 3.5
+ GUIDE_RESCALE:
+ DEFAULT: 0.0
+ VISIBLE: False
+ DISCRETIZATION:
+ VALUES: []
+ DEFAULT:
+ VISIBLE: False
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: bfloat16
+ INPUT: ["LATENT"]
+ PARAS:
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: bfloat16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: bfloat16
+ INPUT: ["PROMPT"]
+#
+MODEL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ NOISE_SCHEDULER:
+ NAME: FlowMatchSigmaScheduler
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ LOGIT_MEAN: 0.0
+ LOGIT_STD: 1.0
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ NAME: FlowMatchFluxShiftScheduler
+ SHIFT: False
+ SIGMOID_SCALE: 1
+ BASE_SHIFT: 0.5
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ NAME: Flux
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
+ IN_CHANNELS: 64
+ HIDDEN_SIZE: 3072
+ NUM_HEADS: 24
+ AXES_DIM: [ 16, 56, 56 ]
+ THETA: 10000
+ VEC_IN_DIM: 768
+ GUIDANCE_EMBED: False
+ CONTEXT_IN_DIM: 4096
+ MLP_RATIO: 4.0
+ QKV_BIAS: True
+ DEPTH: 19
+ DEPTH_SINGLE_BLOCKS: 38
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5PlusClipFluxEmbedder
+ T5_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: T5EncoderModel
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
+ HF_TOKENIZER_CLS: T5Tokenizer
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
+ MAX_LENGTH: 256
+ OUTPUT_KEY: last_hidden_state
+ D_TYPE: bfloat16
+ BATCH_INFER: False
+ CLEAN: whitespace
+ #
+ CLIP_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: CLIPTextModel
+ MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
+ MAX_LENGTH: 77
+ OUTPUT_KEY: pooler_output
+ D_TYPE: bfloat16
+ BATCH_INFER: True
+ CLEAN: whitespace
+#
+MODEL_LOCAL:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ NOISE_SCHEDULER:
+ NAME: FlowMatchSigmaScheduler
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ LOGIT_MEAN: 0.0
+ LOGIT_STD: 1.0
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ NAME: FlowMatchFluxShiftScheduler
+ SHIFT: False
+ SIGMOID_SCALE: 1
+ BASE_SHIFT: 0.5
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ NAME: Flux
+ PRETRAINED_MODEL: models/scepter/FLUX.1-schnell/flux1-schnell.safetensors
+ IN_CHANNELS: 64
+ HIDDEN_SIZE: 3072
+ NUM_HEADS: 24
+ AXES_DIM: [ 16, 56, 56 ]
+ THETA: 10000
+ VEC_IN_DIM: 768
+ GUIDANCE_EMBED: False
+ CONTEXT_IN_DIM: 4096
+ MLP_RATIO: 4.0
+ QKV_BIAS: True
+ DEPTH: 19
+ DEPTH_SINGLE_BLOCKS: 38
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: models/scepter/FLUX.1-schnell/ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5PlusClipFluxEmbedder
+ T5_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: T5EncoderModel
+ MODEL_PATH: models/scepter/FLUX.1-schnell/text_encoder_2/
+ HF_TOKENIZER_CLS: T5Tokenizer
+ TOKENIZER_PATH: models/scepter/FLUX.1-schnell/tokenizer_2/
+ MAX_LENGTH: 256
+ OUTPUT_KEY: last_hidden_state
+ D_TYPE: bfloat16
+ BATCH_INFER: False
+ CLEAN: whitespace
+ #
+ CLIP_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: CLIPTextModel
+ MODEL_PATH: models/scepter/FLUX.1-schnell/text_encoder/
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ TOKENIZER_PATH: models/scepter/FLUX.1-schnell/tokenizer/
+ MAX_LENGTH: 77
+ OUTPUT_KEY: pooler_output
+ D_TYPE: bfloat16
+ BATCH_INFER: True
+ CLEAN: whitespace
+#
+MODEL_HF:
+ NAME: LatentDiffusionFlux
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ DEFAULT_N_PROMPT:
+ USE_EMA: False
+ EVAL_EMA: False
+ DIFFUSION:
+ NAME: DiffusionFluxRF
+ PREDICTION_TYPE: raw
+ NOISE_SCHEDULER:
+ NAME: FlowMatchSigmaScheduler
+ WEIGHTING_SCHEME: logit_normal
+ SHIFT: 3.0
+ LOGIT_MEAN: 0.0
+ LOGIT_STD: 1.0
+ MODE_SCALE: 1.29
+ SAMPLER_SCHEDULER:
+ NAME: FlowMatchFluxShiftScheduler
+ SHIFT: False
+ SIGMOID_SCALE: 1
+ BASE_SHIFT: 0.5
+ MAX_SHIFT: 1.15
+ #
+ DIFFUSION_MODEL:
+ NAME: Flux
+ PRETRAINED_MODEL: hf://black-forest-labs/FLUX.1-schnell@flux1-schnell.safetensors
+ IN_CHANNELS: 64
+ HIDDEN_SIZE: 3072
+ NUM_HEADS: 24
+ AXES_DIM: [ 16, 56, 56 ]
+ THETA: 10000
+ VEC_IN_DIM: 768
+ GUIDANCE_EMBED: False
+ CONTEXT_IN_DIM: 4096
+ MLP_RATIO: 4.0
+ QKV_BIAS: True
+ DEPTH: 19
+ DEPTH_SINGLE_BLOCKS: 38
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKLFlux
+ EMBED_DIM: 16
+ PRETRAINED_MODEL: hf://black-forest-labs/FLUX.1-schnell@ae.safetensors
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 8
+ USE_CONV: False
+ SCALE_FACTOR: 0.3611
+ SHIFT_FACTOR: 0.1159
+ #
+ ENCODER:
+ NAME: Encoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ USE_CHECKPOINT: True
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5PlusClipFluxEmbedder
+ T5_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: T5EncoderModel
+ MODEL_PATH: hf://black-forest-labs/FLUX.1-schnell@text_encoder_2/
+ HF_TOKENIZER_CLS: T5Tokenizer
+ TOKENIZER_PATH: hf://black-forest-labs/FLUX.1-schnell@tokenizer_2/
+ MAX_LENGTH: 256
+ OUTPUT_KEY: last_hidden_state
+ D_TYPE: bfloat16
+ BATCH_INFER: False
+ CLEAN: whitespace
+ #
+ CLIP_MODEL:
+ NAME: HFEmbedder
+ HF_MODEL_CLS: CLIPTextModel
+ MODEL_PATH: hf://black-forest-labs/FLUX.1-schnell@text_encoder/
+ HF_TOKENIZER_CLS: CLIPTokenizer
+ TOKENIZER_PATH: hf://black-forest-labs/FLUX.1-schnell@tokenizer/
+ MAX_LENGTH: 77
+ OUTPUT_KEY: pooler_output
+ D_TYPE: bfloat16
+ BATCH_INFER: True
+ CLEAN: whitespace
\ No newline at end of file
diff --git a/scepter/workflow/config/mantra.yaml b/scepter/workflow/config/mantra.yaml
new file mode 100644
index 0000000..087b20b
--- /dev/null
+++ b/scepter/workflow/config/mantra.yaml
@@ -0,0 +1,1921 @@
+MANTRAS:
+ -
+ NAME: cinematic-diva
+ NAME_ZH: 电影歌星画风
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: UHD, 8K, ultra detailed, a cinematic photograph of {prompt}, beautiful lighting, great composition
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, NSFW
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/967b1852ea26dcf41360fc5542a0df6b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Abstract Expressionism
+ NAME_ZH: 抽象表现主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Abstract Expressionism Art, {prompt}, High contrast, minimalistic, colorful, stark, dramatic, expressionism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/11667c7a085b53a6ffbb76b83788f8b5.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Academia
+ NAME_ZH: 学院风
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Academia, {prompt}, preppy Ivy League style, stark, dramatic, chic boarding school, academia
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, grunge, sloppy, unkempt
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c95197ef78312e2bbd45883bfbfa095a.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Action Figure
+ NAME_ZH: 动作人偶
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Action Figure, {prompt}, plastic collectable action figure, collectable toy action figure
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/24ac67fe21833fc1300fba09d5de6090.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Adorable 3D Character
+ NAME_ZH: 可爱的3D角色
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Adorable 3D Character, {prompt}, 3D render, adorable character, 3D art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, grunge, sloppy, unkempt, photograph, photo, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/31a0ffbfc93b5b336a3bedc0a9985a02.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Adorable Kawaii
+ NAME_ZH: 可爱卡哇伊风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Adorable Kawaii, {prompt}, pretty, cute, adorable, kawaii
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, gothic, dark, moody, monochromatic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/82624c429d504c8290d4c9147be64029.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Art Deco
+ NAME_ZH: 艺术装饰风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Art Deco, {prompt}, sleek, geometric forms, art deco style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c60ef3e7ada8774c92bb6e73c570f819.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Art Nouveau
+ NAME_ZH: 新艺术风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Art Nouveau, beautiful art, {prompt}, sleek, organic forms, long, sinuous, art nouveau style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, industrial, mechanical
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/153d83e71fbf145d0aa5c41ecab0a505.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Astral Aura
+ NAME_ZH: 星体光环
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Astral Aura, {prompt}, astral, colorful aura, vibrant energy
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/e093c9d0f81037002e88f2a6c3f38a55.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Avant-garde
+ NAME_ZH: 先锋派
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Avant-garde, {prompt}, unusual, experimental, avant-garde art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9921d1944893466edfc6a060faf4070e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Baroque
+ NAME_ZH: 巴洛克风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Baroque, {prompt}, dramatic, exuberant, grandeur, baroque art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d8279be4bb34e5cb00a95919412154a4.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Bauhaus-Style Poster
+ NAME_ZH: 包豪斯风格海报
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Bauhaus-Style Poster, {prompt}, simple geometric shapes, clean lines, primary colors, Bauhaus-Style Poster
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/4bd18b32c23e13b9e4fc7f65a1916eab.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Blueprint Schematic Drawing
+ NAME_ZH: 蓝图原理图绘制
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Blueprint Schematic Drawing, {prompt}, technical drawing, blueprint, schematic
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/62392c3715267f7162a37a8ff1caa3a1.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Caricature
+ NAME_ZH: 漫画
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Caricature, {prompt}, exaggerated, comical, caricature
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/894f40ed44b37c3372e6a22b8ae577a4.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Cel Shaded Art
+ NAME_ZH: 单色阴影艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Cel Shaded Art, {prompt}, 2D, flat color, toon shading, cel shaded style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/71c2e8a2cd1031b2bfae1640f7b72c88.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Character Design Sheet
+ NAME_ZH: 角色设计图
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Character Design Sheet, {prompt}, character reference sheet, character turn around
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/eacf2134aa51d4b4a6f2f35cc170f315.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Classicism Art
+ NAME_ZH: 古典主义艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Classicism Art, {prompt}, inspired by Roman and Greek culture, clarity, harmonious, classicism art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/10025b4f3a09e6134086e3cdec4ef2c8.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Color Field Painting
+ NAME_ZH: 色域绘画
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Color Field Painting, {prompt}, abstract, simple, geometic, color field painting style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/80e5b4075c572c04cbb4e48c37b8366b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Colored Pencil Art
+ NAME_ZH: 彩色铅笔艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Colored Pencil Art, {prompt}, colored pencil strokes, light color, visible paper texture, colored pencil art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9ae235d7f1a7c2a4edab52a5e9f9cbae.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Conceptual Art
+ NAME_ZH: 概念艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Conceptual Art, {prompt}, concept art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/cf2e6781997c6842a16155fdff911ef8.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Constructivism
+ NAME_ZH: 结构主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Constructivism Art, {prompt}, minimalistic, geometric forms, constructivism art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d07feccabcadfd3310464bedba858bd1.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Cubism
+ NAME_ZH: 立体主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Cubism Art, {prompt}, flat geometric forms, cubism art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/aa8313f6c3cb9cafb9a7ec07db78df1c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Dadaism
+ NAME_ZH: 达达主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Dadaism Art, {prompt}, satirical, nonsensical, dadaism art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/db27030369cb3ccb042b8aea6bf635e2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Dark Fantasy
+ NAME_ZH: 黑暗幻想
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Dark Fantasy Art, {prompt}, dark, moody, dark fantasy style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, bright, sunny
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c88b6e641c707afc0c8d278ba7e1ac19.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Dark Moody Atmosphere
+ NAME_ZH: 暗色忧郁氛围
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Dark Moody Atmosphere, {prompt}, dramatic, mysterious, dark moody atmosphere
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, vibrant, colorful, bright
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/3da915da2f5cedaf243e57e08163f35b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: DMT Art Style
+ NAME_ZH: DMT艺术风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: DMT Art Style, {prompt}, bright colors, surreal visuals, swirling patterns, DMT art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d4a823d2bfa912ca4bb56d96f722f08f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Doodle Art
+ NAME_ZH: 涂鸦艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Doodle Art Style, {prompt}, drawing, freeform, swirling patterns, doodle art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8e1e21745c149b9634d3ce96fc7d505f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Double Exposure
+ NAME_ZH: 双重曝光
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Double Exposure Style, {prompt}, double image ghost effect, image combination, double exposure style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/efd7aea4af1c4ede99fdfac17350e264.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Dripping Paint Splatter Art
+ NAME_ZH: 滴漆溅画艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Dripping Paint Splatter Art, {prompt}, dramatic, paint drips, splatters, dripping paint
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/69fd81f5983107acc3d334af62915851.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Expressionism
+ NAME_ZH: 表现主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Expressionism Art Style, {prompt}, movement, contrast, emotional, exaggerated forms, expressionism art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/3dfa862bd1c80cec237e4a5717cea2bd.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Faded Polaroid Photo
+ NAME_ZH: 褪色的宝丽来照片
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Faded Polaroid Photo, {prompt}, analog, old faded photo, old polaroid
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, vibrant, colorful
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/f152edb4b3ca6248758b48115258ddfa.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fauvism
+ NAME_ZH: 野兽派
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Fauvism Art, {prompt}, painterly, bold colors, textured brushwork, fauvism art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0506b350cadcf8fca42da764cd6fd5bf.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Flat 2D Art
+ NAME_ZH: 扁平2D艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Flat 2D Art, {prompt}, simple flat color, 2-dimensional, Flat 2D Art Style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, 3D, photo, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/940cfd34155634cf051e1b2942cca426.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fortnite Art Style
+ NAME_ZH: 堡垒之夜艺术风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Fortnite Art Style, {prompt}, 3D cartoon, colorful, Fortnite Art Style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, photo, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/156f8d943ff6d283f7a34f265daaa46c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Futurism
+ NAME_ZH: 未来主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Futurism Art Style, {prompt}, dynamic, dramatic, Futurism Art Style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/855283cc2ab6283b627ef32d06a1ae0f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Glitchcore
+ NAME_ZH: 故障核心
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Glitchcore Art Style, {prompt}, dynamic, dramatic, distorted, vibrant colors, glitchcore art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9af62376402a85774179e82cf7e0dc59.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Glo-fi
+ NAME_ZH: 光环音乐风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Glo-fi Art Style, {prompt}, dynamic, dramatic, vibrant colors, glo-fi art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/275a6e2a3297cec57f9b5a4a17b0f749.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Googie Art Style
+ NAME_ZH: 古奇艺术风格
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Googie Art Style, {prompt}, dynamic, dramatic, 1950's futurism, bold boomerang angles, Googie art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/08ee77586d7bec3f6bad7d60cee3c540.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Graffiti Art
+ NAME_ZH: 涂鸦艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Graffiti Art Style, {prompt}, dynamic, dramatic, vibrant colors, graffiti art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/57b751b11564cb22cd49ef21f2004a5f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Harlem Renaissance Art
+ NAME_ZH: 哈莱姆文艺复兴艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Harlem Renaissance Art Style, {prompt}, dynamic, dramatic, 1920s African American culture, Harlem Renaissance art style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d09ea3dadcd0a15d8fa247ee74a75e6b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: High Fashion
+ NAME_ZH: 高级时装
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: High Fashion, {prompt}, dynamic, dramatic, haute couture, elegant, ornate clothing, High Fashion
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/868f752bbb5ef992a0be36d7d13de3be.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Idyllic
+ NAME_ZH: 田园诗般的
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Idyllic, {prompt}, peaceful, happy, pleasant, happy, harmonious, picturesque, charming
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/55bab29eeb628e7a9ae018e45e7e24db.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Impressionism
+ NAME_ZH: 印象主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Impressionism, {prompt}, painterly, small brushstrokes, visible brushstrokes, impressionistic style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0312b673dc6858a9864d7f45f0c5c1fc.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Infographic Drawing
+ NAME_ZH: 信息图表绘制
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Infographic Drawing, {prompt}, diagram, infographic
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/b09151f6e0883d26056b80fcc9d398bb.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Ink Dripping Drawing
+ NAME_ZH: 墨水滴画
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Ink Dripping Drawing, {prompt}, ink drawing, dripping ink
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, colorful, vibrant
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/5f09912a07e915250a96b2023db15821.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Japanese Ink Drawing
+ NAME_ZH: 日本墨画
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Japanese Ink Drawing, {prompt}, ink drawing, inkwash, Japanese Ink Drawing
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, colorful, vibrant
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/633ad971ff30fe4717782fc6b5f92fda.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Knolling Photography
+ NAME_ZH: 秩序拍摄
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Knolling Photography, {prompt}, flat lay photography, object arrangment, knolling photography
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/89a205be3a276349ecddf0eeca43ed80.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Light Cheery Atmosphere
+ NAME_ZH: 轻快愉快的氛围
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Light Cheery Atmosphere, {prompt}, happy, joyful, cheerful, carefree, gleeful, lighthearted, pleasant atmosphere
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, monochromatic, dark, moody
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0913ac9815f892411f8327a630f51ae4.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Logo Design
+ NAME_ZH: 标志设计
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Logo Design, {prompt}, dynamic graphic art, vector art, minimalist, professional logo design
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9aa040b0c60d289da9610c91ad9b7c7e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Luxurious Elegance
+ NAME_ZH: 奢华优雅
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Luxurious Elegance, {prompt}, extravagant, ornate, designer, opulent, picturesque, lavish
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/e796e84ed745150f0d4e28e0ac99b4cf.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Macro Photography
+ NAME_ZH: 微距摄影
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Macro Photography, {prompt}, close-up, macro 100mm, macro photography
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/92e8b84379828f38e3d01c7272f41b0b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Mandola Art
+ NAME_ZH: 曼陀罗艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Mandola art style, {prompt}, complex, circular design, mandola
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/dbf5ca944d9213c3181348666cf337ac.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Marker Drawing
+ NAME_ZH: 马克笔绘图
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Marker Drawing, {prompt}, bold marker lines, visibile paper texture, marker drawing
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, photograph, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c558745f7a4b77d7ca9a428d8874efe4.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Medievalism
+ NAME_ZH: 中世纪主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Medievalism, {prompt}, inspired by The Middle Ages, medieval art, elaborate patterns and decoration, Medievalism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/26aa6244359a4fd87f6438a341057901.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Minimalism
+ NAME_ZH: 极简主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Minimalism, {prompt}, abstract, simple geometic shapes, hard edges, sleek contours, Minimalism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/dce92a90da6299339cf8e9ebc6596ea2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Neo-Baroque
+ NAME_ZH: 新巴洛克
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Neo-Baroque, {prompt}, ornate and elaborate, dynaimc, Neo-Baroque
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9b6cd8751b18c65b259769b44cfac351.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Neo-Byzantine
+ NAME_ZH: 新拜占庭
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Neo-Byzantine, {prompt}, grand decorative religious style, Orthodox Christian inspired, Neo-Byzantine
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/99c67768869b546bf527ef0b1735c0d5.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Neo-Futurism
+ NAME_ZH: 新未来主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Neo-Futurism, {prompt}, high-tech, curves, spirals, flowing lines, idealistic future, Neo-Futurism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/bada1534f7a60187f584febcc92f40d1.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Neo-Impressionism
+ NAME_ZH: 新印象主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Neo-Impressionism, {prompt}, tiny dabs of color, Pointillism, painterly, Neo-Impressionism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, photograph, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/6bafeac67ce2679b64a86ee1023b53c9.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Neo-Rococo
+ NAME_ZH: 新洛可可
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Neo-Rococo, {prompt}, curved forms, naturalistic ornamentation, elaborate, decorative, gaudy, Neo-Rococo
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9ab0ad2b88e03933ea479357ff1e4435.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Neoclassicism
+ NAME_ZH: 新古典主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Neoclassicism, {prompt}, ancient Rome and Greece inspired, idealic, sober colors, Neoclassicism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/83a4e3a2ca577400b05a81b5b2b8650e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Op Art
+ NAME_ZH: 视觉艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Op Art, {prompt}, optical illusion, abstract, geometric pattern, impression of movement, Op Art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/5f11132803dff2de5213293b91d407d7.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Ornate and Intricate
+ NAME_ZH: 华丽复杂
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Ornate and Intricate, {prompt}, decorative, highly detailed, elaborate, ornate, intricate
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/24b895dca946c8765ad9fa38f720f671.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Pencil Sketch Drawing
+ NAME_ZH: 铅笔素描
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Pencil Sketch Drawing, {prompt}, black and white drawing, graphite drawing
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a9056e1eac85e5e4fe96a93917d4cce4.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Pop Art 2
+ NAME_ZH: 流行艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Pop Art, {prompt}, vivid colors, flat color, 2D, strong lines, Pop Art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, photo, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/83c4c06c70a48da6558ff970ece247b6.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Rococo
+ NAME_ZH: 洛可可
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Rococo, {prompt}, flamboyant, pastel colors, curved lines, elaborate detail, Rococo
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/4850170f410e10bc833b7d00324dff18.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Silhouette Art
+ NAME_ZH: 剪影艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Silhouette Art, {prompt}, high contrast, well defined, Silhouette Art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/568777f447fc02510b618152726d5002.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Simple Vector Art
+ NAME_ZH: 简单矢量艺术
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Simple Vector Art, {prompt}, 2D flat, simple shapes, minimalistic, professional graphic, flat color, high contrast, Simple Vector Art
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, 3D, photo, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/24d2ab1fd175bf4e9ef3a8327651dd4b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Sketchup
+ NAME_ZH: 草图大师
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Sketchup, {prompt}, CAD, professional design, Sketchup
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, photo, photograph
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/337dca7ac49cf7820c85ede096ebde9c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Steampunk 2
+ NAME_ZH: 蒸汽朋克
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Steampunk, {prompt}, retrofuturistic science fantasy, steam-powered tech, vintage industry, gears, neo-victorian, steampunk
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/07d7b27cd73f2d43684003563511c15b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Surrealism
+ NAME_ZH: 超现实主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Surrealism, {prompt}, expressive, dramatic, organic lines and forms, dreamlike and mysterious, Surrealism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/38bae70377b323d1a5f5756d7e04886d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Suprematism
+ NAME_ZH: 至上主义
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Suprematism, {prompt}, abstract, limited color palette, geometric forms, Suprematism
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c48fbd8cedee84d06b4d875d566f1bb8.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Terragen
+ NAME_ZH: 地形生成
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Terragen, {prompt}, beautiful massive landscape, epic scenery, Terragen
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8974cd8e0f38fb7aac117cf0983fc36d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Tranquil Relaxing Atmosphere
+ NAME_ZH: 宁静放松的氛围
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Tranquil Relaxing Atmosphere, {prompt}, calming style, soothing colors, peaceful, idealic, Tranquil Relaxing Atmosphere
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, oversaturated
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/69170d8443a67be4210af8d5558779c2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Sticker Designs
+ NAME_ZH: 贴纸设计
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Vector Art Stickers, {prompt}, professional vector design, sticker designs, Sticker Sheet
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2d1e9867058db2c57f2fe47530de3243.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Vibrant Rim Light
+ NAME_ZH: 生动的边缘光
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Vibrant Rim Light, {prompt}, bright rim light, high contrast, bold edge light
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/bb08936ca184daba2b30b4a9efb308f2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Volumetric Lighting
+ NAME_ZH: 体积光照明
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Volumetric Lighting, {prompt}, light depth, dramatic atmospheric lighting, Volumetric Lighting
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/bbad7964a2ae70b66c72804273f11e74.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Watercolor 2
+ NAME_ZH: 水彩
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Watercolor style painting, {prompt}, visible paper texture, colorwash, watercolor
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, photo, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8859d532ae5901cc8457d6118fb9b7da.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Whimsical and Playful
+ NAME_ZH: 异想天开和俏皮
+ DESCRIPTION:
+ SOURCE: diva
+ PROMPT: Whimsical and Playful, {prompt}, imaginative, fantastical, bight colors, stylized, happy, Whimsical and Playful
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, drab, boring, moody
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/af2267f1942e870be957c37cd73d4359.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fooocus Enhance
+ NAME_ZH: 焦点增强
+ DESCRIPTION:
+ SOURCE: fooocus
+ PROMPT: {
+ "prompt": null
+}
+
+ NEGATIVE_PROMPT: (worst quality, low quality, normal quality, lowres, low details, oversaturated, undersaturated, overexposed, underexposed, grayscale, bw, bad photo, bad photography, bad art:1.4), (watermark, signature, text font, username, error, logo, words, letters, digits, autograph, trademark, name:1.2), (blur, blurry, grainy), morbid, ugly, asymmetrical, mutated malformed, mutilated, poorly lit, bad shadow, draft, cropped, out of frame, cut off, censored, jpeg artifacts, out of focus, glitch, duplicate, (airbrushed, cartoon, anime, semi-realistic, cgi, render, blender, digital art, manga, amateur:1.3), (3D ,3D Game, 3D Game Scene, 3D Character:1.1), (bad hands, bad anatomy, bad body, bad face, bad teeth, bad arms, bad legs, deformities:1.3)
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/e501823a21bbda56592055c9613c1dbb.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fooocus Sharp
+ NAME_ZH: 焦点锐化
+ DESCRIPTION:
+ SOURCE: fooocus
+ PROMPT: cinematic still {prompt} . emotional, harmonious, vignette, 4k epic detailed, shot on kodak, 35mm photo, sharp focus, high budget, cinemascope, moody, epic, gorgeous, film grain, grainy
+ NEGATIVE_PROMPT: anime, cartoon, graphic, (blur, blurry, bokeh), text, painting, crayon, graphite, abstract, glitch, deformed, mutated, ugly, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/393fb7aab57992a1c3af7cbad01c1001.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fooocus Masterpiece
+ NAME_ZH: 焦点杰作
+ DESCRIPTION:
+ SOURCE: fooocus
+ PROMPT: (masterpiece), (best quality), (ultra-detailed), {prompt}, illustration, disheveled hair, detailed eyes, perfect composition, moist skin, intricate details, earrings, by wlop
+ NEGATIVE_PROMPT: longbody, lowres, bad anatomy, bad hands, missing fingers, pubic hair,extra digit, fewer digits, cropped, worst quality, low quality
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9e27c3d1475dacae0ee45cea783aaaf5.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fooocus Photograph
+ NAME_ZH: 焦点摄影
+ DESCRIPTION:
+ SOURCE: fooocus
+ PROMPT: photograph {prompt}, 50mm . cinematic 4k epic detailed 4k epic detailed photograph shot on kodak detailed cinematic hbo dark moody, 35mm photo, grainy, vignette, vintage, Kodachrome, Lomography, stained, highly detailed, found footage
+ NEGATIVE_PROMPT: Brad Pitt, bokeh, depth of field, blurry, cropped, regular face, saturated, contrast, deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, cartoon, drawing, anime, text, cropped, out of frame, worst quality, low quality, jpeg artifacts, ugly, duplicate, morbid, mutilated, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, deformed, dehydrated, bad anatomy, bad proportions, extra limbs, cloned face, disfigured, gross proportions, malformed limbs, missing arms, missing legs, extra arms, extra legs, fused fingers, too many fingers, long neck
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8b2ea41ec15bf1a16d6a837d46dbc2ac.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fooocus Negative
+ NAME_ZH: 焦点底片
+ DESCRIPTION:
+ SOURCE: fooocus
+ PROMPT:
+ NEGATIVE_PROMPT: deformed, bad anatomy, disfigured, poorly drawn face, mutated, extra limb, ugly, poorly drawn hands, missing limb, floating limbs, disconnected limbs, disconnected head, malformed hands, long neck, mutated hands and fingers, bad hands, missing fingers, cropped, worst quality, low quality, mutation, poorly drawn, huge calf, bad hands, fused hand, missing hand, disappearing arms, disappearing thigh, disappearing calf, disappearing legs, missing fingers, fused fingers, abnormal eye proportion, Abnormal hands, abnormal legs, abnormal feet, abnormal fingers, drawing, painting, crayon, sketch, graphite, impressionist, noisy, blurry, soft, deformed, ugly, anime, cartoon, graphic, text, painting, crayon, graphite, abstract, glitch
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/daaec2962291137532189b8a31012532.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: Fooocus Cinematic
+ NAME_ZH: 焦点电影
+ DESCRIPTION:
+ SOURCE: fooocus
+ PROMPT: cinematic still {prompt} . emotional, harmonious, vignette, highly detailed, high budget, bokeh, cinemascope, moody, epic, gorgeous, film grain, grainy
+ NEGATIVE_PROMPT: anime, cartoon, graphic, text, painting, crayon, graphite, abstract, glitch, deformed, mutated, ugly, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/adc4d52aa5b0afca0593a475ddc9d055.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-cinematic-dynamic
+ NAME_ZH: MRE电影动态
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: epic cinematic shot of dynamic {prompt} in motion. main subject of high budget action movie. raw photo, motion blur. best quality, high resolution
+ NEGATIVE_PROMPT: static, still, motionless, sluggish. drawing, painting, illustration, rendered. low budget. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/bd2470068e09d9f1b7d0a690c879b3ce.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-spontaneous-picture
+ NAME_ZH: MRE自发图片
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: spontaneous picture of {prompt}, taken by talented amateur. best quality, high resolution. magical moment, natural look. simple but good looking
+ NEGATIVE_PROMPT: overthinked. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/19ad6082cee5516ce330641b515467ef.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-artistic-vision
+ NAME_ZH: MRE艺术视觉
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: powerful artistic vision of {prompt}. breathtaking masterpiece made by great artist. best quality, high resolution
+ NEGATIVE_PROMPT: insignificant, flawed, made by bad artist. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/95af94e50bcb91ae8fd2914a3bf13f1e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-dark-dream
+ NAME_ZH: MRE黑暗梦境
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: dark and unsettling dream showing {prompt}. best quality, high resolution. created by genius but depressed mad artist. grim beauty
+ NEGATIVE_PROMPT: naive, cheerful. comfortable, casual, boring, cliche. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/572c8c26cc20ac0ee66684e3c5ee4e8c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-gloomy-art
+ NAME_ZH: MRE忧郁艺术
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: astonishing gloomy art made mainly of shadows and lighting, forming {prompt}. masterful usage of lighting, shadows and chiaroscuro. made by black-hearted artist, drawing from darkness. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/66cbdf819e930f7580bd66a41bde7dfe.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-bad-dream
+ NAME_ZH: MRE恶梦
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: picture from really bad dream about terrifying {prompt}, true horror. bone-chilling vision. mad world that shouldn't exist. best quality, high resolution
+ NEGATIVE_PROMPT: nice dream, pleasant experience. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/b92b2bbfc400db9204fe8249354132b4.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-underground
+ NAME_ZH: MRE地下
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: uncanny caliginous vision of {prompt}, created by remarkable underground artist. best quality, high resolution. raw and brutal art, careless but impressive style. inspired by darkness and chaos
+ NEGATIVE_PROMPT: photography, mainstream, civilized. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8d46eca558d791a1f2b41b0ed7bba4d2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-surreal-painting
+ NAME_ZH: MRE超现实绘画
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: surreal painting representing strange vision of {prompt}. harmonious madness, synergy with chance. unique artstyle, mindbending art, magical surrealism. best quality, high resolution
+ NEGATIVE_PROMPT: photography, illustration, drawing. realistic, possible. logical, sane. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/58858425832b10d233f7887af4a2022f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-dynamic-illustration
+ NAME_ZH: MRE动态插画
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: insanely dynamic illustration of {prompt}. best quality, high resolution. crazy artstyle, careless brushstrokes, emotional and fun
+ NEGATIVE_PROMPT: photography, realistic. static, still, slow, boring. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8bf9861d4d3fadcdb98ceec88a941582.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-undead-art
+ NAME_ZH: MRE不死艺术
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: long forgotten art created by undead artist illustrating {prompt}, tribute to the death and decay. miserable art of the damned. wretched and decaying world. best quality, high resolution
+ NEGATIVE_PROMPT: alive, playful, living. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/fada50979ca180006eba9a45a00f3675.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-elemental-art
+ NAME_ZH: MRE元素艺术
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: art illustrating insane amounts of raging elemental energy turning into {prompt}, avatar of elements. magical surrealism, wizardry. best quality, high resolution
+ NEGATIVE_PROMPT: photography, realistic, real. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/5895d78cf58c1ca05178991f37cc48ff.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-space-art
+ NAME_ZH: MRE太空艺术
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: winner of inter-galactic art contest illustrating {prompt}, symbol of the interstellar singularity. best quality, high resolution. artstyle previously unseen in the whole galaxy
+ NEGATIVE_PROMPT: created by human race, low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/e9815495587895a21d970728474c8be6.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-ancient-illustration
+ NAME_ZH: MRE古代插画
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: sublime ancient illustration of {prompt}, predating human civilization. crude and simple, but also surprisingly beautiful artwork, made by genius primeval artist. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/151e07c17a89ebaf7688905ea0199862.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-brave-art
+ NAME_ZH: MRE勇敢艺术
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: brave, shocking, and brutally true art showing {prompt}. inspired by courage and unlimited creativity. truth found in chaos. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/1a068f2728327d21ebb99285c2ce1370.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-heroic-fantasy
+ NAME_ZH: MRE英雄幻想
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: heroic fantasy painting of {prompt}, in the dangerous fantasy world. airbrush over oil on canvas. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/78d05472067e0b86c8270fa5476fcb8e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-dark-cyberpunk
+ NAME_ZH: MRE黑暗赛博朋克
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: dark cyberpunk illustration of brutal {prompt} in a world without hope, ruled by ruthless criminal corporations. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a9f50e3162958fd783872d6955ad7d0b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-lyrical-geometry
+ NAME_ZH: MRE抒情几何
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: geometric and lyrical abstraction painting presenting {prompt}. oil on metal. best quality, high resolution
+ NEGATIVE_PROMPT: photography, realistic, drawing, rendered. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/dfabd41d3042ced804bc97ae35e3a7cb.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-sumi-e-symbolic
+ NAME_ZH: MRE墨绘象征
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: big long brushstrokes of deep black sumi-e turning into symbolic painting of {prompt}. master level raw art. best quality, high resolution
+ NEGATIVE_PROMPT: photography, rendered. low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/695b1ba687544eaeec9fb5f871917aeb.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-sumi-e-detailed
+ NAME_ZH: MRE墨绘精细
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: highly detailed black sumi-e painting of {prompt}. in-depth study of perfection, created by a master. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/16eb95e180385b88794e368b28da1812.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-manga
+ NAME_ZH: MRE漫画
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: manga artwork presenting {prompt}. created by japanese manga artist. highly emotional. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/ebec631bf467937f82d05958ae59f9dc.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-anime
+ NAME_ZH: MRE动漫
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: anime artwork illustrating {prompt}. created by japanese anime studio. highly emotional. best quality, high resolution
+ NEGATIVE_PROMPT: low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a08149bc8e50f6bc65c0010d4cd416f8.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: mre-comic
+ NAME_ZH: MRE漫画书
+ DESCRIPTION:
+ SOURCE: mre
+ PROMPT: breathtaking illustration from adult comic book presenting {prompt}. fabulous artwork. best quality, high resolution
+ NEGATIVE_PROMPT: deformed, ugly, low quality, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/48c65cebf1fa4284d7b8feb619412e65.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-3d-model
+ NAME_ZH: SAI三维模型
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: professional 3d model {prompt} . octane render, highly detailed, volumetric, dramatic lighting
+ NEGATIVE_PROMPT: ugly, deformed, noisy, low poly, blurry, painting
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c1a765ac089fdfb3c1d11b33c75d2afd.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-analog film
+ NAME_ZH: SAI模拟胶片
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: analog film photo {prompt} . faded film, desaturated, 35mm photo, grainy, vignette, vintage, Kodachrome, Lomography, stained, highly detailed, found footage
+ NEGATIVE_PROMPT: painting, drawing, illustration, glitch, deformed, mutated, cross-eyed, ugly, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/3c46d19957efd7fb78f4ab2bdada5468.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-anime
+ NAME_ZH: SAI动漫
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: anime artwork {prompt} . anime style, key visual, vibrant, studio anime, highly detailed
+ NEGATIVE_PROMPT: photo, deformed, black and white, realism, disfigured, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c030e72561eda96abcf738f1370d36ff.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-cinematic
+ NAME_ZH: SAI电影
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: cinematic film still {prompt} . shallow depth of field, vignette, highly detailed, high budget, bokeh, cinemascope, moody, epic, gorgeous, film grain, grainy
+ NEGATIVE_PROMPT: anime, cartoon, graphic, text, painting, crayon, graphite, abstract, glitch, deformed, mutated, ugly, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d9161c0d5cbf2133b2bfc1021c0c5e2a.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-comic book
+ NAME_ZH: SAI漫画书
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: comic {prompt} . graphic illustration, comic art, graphic novel art, vibrant, highly detailed
+ NEGATIVE_PROMPT: photograph, deformed, glitch, noisy, realistic, stock photo
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/6759ecb831037367e64c4b36d802be87.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-craft clay
+ NAME_ZH: SAI手工粘土
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: play-doh style {prompt} . sculpture, clay art, centered composition, Claymation
+ NEGATIVE_PROMPT: sloppy, messy, grainy, highly detailed, ultra textured, photo
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8fc51113f725f27326c4398a7457cd6d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-digital art
+ NAME_ZH: SAI数字艺术
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: concept art {prompt} . digital artwork, illustrative, painterly, matte painting, highly detailed
+ NEGATIVE_PROMPT: photo, photorealistic, realism, ugly
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d29adb44458700c4a45ee6edaa04bfb6.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-enhance
+ NAME_ZH: SAI增强
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: breathtaking {prompt} . award-winning, professional, highly detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, distorted, grainy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2be0de541ff65f8da80ddc24a65c98d8.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-fantasy art
+ NAME_ZH: SAI幻想艺术
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: ethereal fantasy concept art of {prompt} . magnificent, celestial, ethereal, painterly, epic, majestic, magical, fantasy art, cover art, dreamy
+ NEGATIVE_PROMPT: photographic, realistic, realism, 35mm film, dslr, cropped, frame, text, deformed, glitch, noise, noisy, off-center, deformed, cross-eyed, closed eyes, bad anatomy, ugly, disfigured, sloppy, duplicate, mutated, black and white
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a6f8d92afcd5803dfb2ebecbc92091b6.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-isometric
+ NAME_ZH: SAI等距
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: isometric style {prompt} . vibrant, beautiful, crisp, detailed, ultra detailed, intricate
+ NEGATIVE_PROMPT: deformed, mutated, ugly, disfigured, blur, blurry, noise, noisy, realistic, photographic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d4e3fcbbfd7b1323bd89decf5d7b0006.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-line art
+ NAME_ZH: SAI线条艺术
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: line art drawing {prompt} . professional, sleek, modern, minimalist, graphic, line art, vector graphics
+ NEGATIVE_PROMPT: anime, photorealistic, 35mm film, deformed, glitch, blurry, noisy, off-center, deformed, cross-eyed, closed eyes, bad anatomy, ugly, disfigured, mutated, realism, realistic, impressionism, expressionism, oil, acrylic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/034a51b0dd34b018be8859bf45b4f7ed.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-lowpoly
+ NAME_ZH: SAI低多边形
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: low-poly style {prompt} . low-poly game art, polygon mesh, jagged, blocky, wireframe edges, centered composition
+ NEGATIVE_PROMPT: noisy, sloppy, messy, grainy, highly detailed, ultra textured, photo
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/76b9913e9fa5704b6d30adbde9e1f70f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-neonpunk
+ NAME_ZH: SAI霓虹朋克
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: neonpunk style {prompt} . cyberpunk, vaporwave, neon, vibes, vibrant, stunningly beautiful, crisp, detailed, sleek, ultramodern, magenta highlights, dark purple shadows, high contrast, cinematic, ultra detailed, intricate, professional
+ NEGATIVE_PROMPT: painting, drawing, illustration, glitch, deformed, mutated, cross-eyed, ugly, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/7e9ed25bb34008beb5f417df63c4b2fe.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-origami
+ NAME_ZH: SAI折纸
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: origami style {prompt} . paper art, pleated paper, folded, origami art, pleats, cut and fold, centered composition
+ NEGATIVE_PROMPT: noisy, sloppy, messy, grainy, highly detailed, ultra textured, photo
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/924f46a8f276011a0953d7988e90ee25.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-photographic
+ NAME_ZH: SAI摄影
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: cinematic photo {prompt} . 35mm photograph, film, bokeh, professional, 4k, highly detailed
+ NEGATIVE_PROMPT: drawing, painting, crayon, sketch, graphite, impressionist, noisy, blurry, soft, deformed, ugly
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d6a2d8f3d37cc21c20c5dfc13d000b67.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-pixel art
+ NAME_ZH: SAI像素艺术
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: pixel-art {prompt} . low-res, blocky, pixel art style, 8-bit graphics
+ NEGATIVE_PROMPT: sloppy, messy, blurry, noisy, highly detailed, ultra textured, photo, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a5ab89c0960be8c1216e65c98d92ae4a.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: sai-texture
+ NAME_ZH: SAI质地
+ DESCRIPTION:
+ SOURCE: sai
+ PROMPT: texture {prompt} top down close-up
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/ebfeab574283fff2ac096e348a585e6d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-advertising
+ NAME_ZH: 广告
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: advertising poster style {prompt} . Professional, modern, product-focused, commercial, eye-catching, highly detailed
+ NEGATIVE_PROMPT: noisy, blurry, amateurish, sloppy, unattractive
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/1aa26a16126e2756bf4bf3fda29baa12.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-automotive
+ NAME_ZH: 汽车广告
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: automotive advertisement style {prompt} . sleek, dynamic, professional, commercial, vehicle-focused, high-resolution, highly detailed
+ NEGATIVE_PROMPT: noisy, blurry, unattractive, sloppy, unprofessional
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/ddc833b36ca23c85a4f6e7bf0088d081.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-corporate
+ NAME_ZH: 企业广告
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: corporate branding style {prompt} . professional, clean, modern, sleek, minimalist, business-oriented, highly detailed
+ NEGATIVE_PROMPT: noisy, blurry, grungy, sloppy, cluttered, disorganized
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/df97fa45c7842296c138aca4057db272.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-fashion editorial
+ NAME_ZH: 时尚编辑
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: fashion editorial style {prompt} . high fashion, trendy, stylish, editorial, magazine style, professional, highly detailed
+ NEGATIVE_PROMPT: outdated, blurry, noisy, unattractive, sloppy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/89259be723479cd0b81546b715bea04d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-food photography
+ NAME_ZH: 食品摄影
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: food photography style {prompt} . appetizing, professional, culinary, high-resolution, commercial, highly detailed
+ NEGATIVE_PROMPT: unappetizing, sloppy, unprofessional, noisy, blurry
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/33acc98f8615940c8cebbac8ce5297b2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-gourmet food photography
+ NAME_ZH: 美食摄影
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: gourmet food photo of {prompt} . soft natural lighting, macro details, vibrant colors, fresh ingredients, glistening textures, bokeh background, styled plating, wooden tabletop, garnished, tantalizing, editorial quality
+ NEGATIVE_PROMPT: cartoon, anime, sketch, grayscale, dull, overexposed, cluttered, messy plate, deformed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/897bbd3f1d232de122796266f3845b50.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-luxury
+ NAME_ZH: 奢华广告
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: luxury product style {prompt} . elegant, sophisticated, high-end, luxurious, professional, highly detailed
+ NEGATIVE_PROMPT: cheap, noisy, blurry, unattractive, amateurish
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/5b33ecc84dc5ff285a0f2e69f43162cb.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-real estate
+ NAME_ZH: 房地产广告
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: real estate photography style {prompt} . professional, inviting, well-lit, high-resolution, property-focused, commercial, highly detailed
+ NEGATIVE_PROMPT: dark, blurry, unappealing, noisy, unprofessional
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/1b99b7bcd0476144d80d066e747d8dfc.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: ads-retail
+ NAME_ZH: 零售广告
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: retail packaging style {prompt} . vibrant, enticing, commercial, product-focused, eye-catching, professional, highly detailed
+ NEGATIVE_PROMPT: noisy, blurry, amateurish, sloppy, unattractive
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0029e46ef36380dde2084257108f8a6f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-abstract
+ NAME_ZH: 抽象艺术风格
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: abstract style {prompt} . non-representational, colors and shapes, expression of feelings, imaginative, highly detailed
+ NEGATIVE_PROMPT: realistic, photographic, figurative, concrete
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8c401cb0a6ea288230222c4985e78667.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-abstract expressionism
+ NAME_ZH: 抽象表现主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: abstract expressionist painting {prompt} . energetic brushwork, bold colors, abstract forms, expressive, emotional
+ NEGATIVE_PROMPT: realistic, photorealistic, low contrast, plain, simple, monochrome
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9df6635149c576b65d911c2e4cc6a86b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-art deco
+ NAME_ZH: 艺术装饰风格
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: art deco style {prompt} . geometric shapes, bold colors, luxurious, elegant, decorative, symmetrical, ornate, detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, modernist, minimalist
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/82fff8a25de8075777caf14cb3eb5650.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-art nouveau
+ NAME_ZH: 新艺术风格
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: art nouveau style {prompt} . elegant, decorative, curvilinear forms, nature-inspired, ornate, detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, modernist, minimalist
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/84075c5c0cb4b7a2541c6fbca9835cfd.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-constructivist
+ NAME_ZH: 构成主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: constructivist style {prompt} . geometric shapes, bold colors, dynamic composition, propaganda art style
+ NEGATIVE_PROMPT: realistic, photorealistic, low contrast, plain, simple, abstract expressionism
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2a634ad0c89ecefbbea4ce5bcab3d5e5.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-cubist
+ NAME_ZH: 立体主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: cubist artwork {prompt} . geometric shapes, abstract, innovative, revolutionary
+ NEGATIVE_PROMPT: anime, photorealistic, 35mm film, deformed, glitch, low contrast, noisy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2a3d797008c08e12b485d61624741ea6.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-expressionist
+ NAME_ZH: 表现主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: expressionist {prompt} . raw, emotional, dynamic, distortion for emotional effect, vibrant, use of unusual colors, detailed
+ NEGATIVE_PROMPT: realism, symmetry, quiet, calm, photo
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/83914655659001716a1acede295289d0.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-graffiti
+ NAME_ZH: 涂鸦
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: graffiti style {prompt} . street art, vibrant, urban, detailed, tag, mural
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/f72a23c7623eb480737eb73c08bf8423.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-hyperrealism
+ NAME_ZH: 超现实主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: hyperrealistic art {prompt} . extremely high-resolution details, photographic, realism pushed to extreme, fine texture, incredibly lifelike
+ NEGATIVE_PROMPT: simplified, abstract, unrealistic, impressionistic, low resolution
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/db378d5a64a8e7ad1e29a4389bda5d1c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-impressionist
+ NAME_ZH: 印象主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: impressionist painting {prompt} . loose brushwork, vibrant color, light and shadow play, captures feeling over form
+ NEGATIVE_PROMPT: anime, photorealistic, 35mm film, deformed, glitch, low contrast, noisy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/53b151aec4d5685dfc24511b6705b90e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-pointillism
+ NAME_ZH: 点彩主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: pointillism style {prompt} . composed entirely of small, distinct dots of color, vibrant, highly detailed
+ NEGATIVE_PROMPT: line drawing, smooth shading, large color fields, simplistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/23b137a409ee8a8c6ee160c1ddf0659f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-pop art
+ NAME_ZH: 波普艺术
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: pop Art style {prompt} . bright colors, bold outlines, popular culture themes, ironic or kitsch
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, minimalist
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/bd4faf0e2b7dbc2d0eb21f3ebee97d3d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-psychedelic
+ NAME_ZH: 迷幻艺术
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: psychedelic style {prompt} . vibrant colors, swirling patterns, abstract forms, surreal, trippy
+ NEGATIVE_PROMPT: monochrome, black and white, low contrast, realistic, photorealistic, plain, simple
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/7dc0131817f4c31517581dd4a811067b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-renaissance
+ NAME_ZH: 文艺复兴
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: renaissance style {prompt} . realistic, perspective, light and shadow, religious or mythological themes, highly detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, modernist, minimalist, abstract
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/69796484f61dc2e94d5853f5bbe27c05.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-steampunk
+ NAME_ZH: 蒸汽朋克
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: steampunk style {prompt} . antique, mechanical, brass and copper tones, gears, intricate, detailed
+ NEGATIVE_PROMPT: deformed, glitch, noisy, low contrast, anime, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/985005e28858b03c9500a6eb5bd9e201.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-surrealist
+ NAME_ZH: 超现实主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: surrealist art {prompt} . dreamlike, mysterious, provocative, symbolic, intricate, detailed
+ NEGATIVE_PROMPT: anime, photorealistic, realistic, deformed, glitch, noisy, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/892d19ec3c429b7148562796df0024db.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-typography
+ NAME_ZH: 排版艺术
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: typographic art {prompt} . stylized, intricate, detailed, artistic, text-based
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9adafe1bb169e1ee3f78ba2c1b1cf8d3.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: artstyle-watercolor
+ NAME_ZH: 水彩艺术
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: watercolor painting {prompt} . vibrant, beautiful, painterly, detailed, textural, artistic
+ NEGATIVE_PROMPT: anime, photorealistic, 35mm film, deformed, glitch, low contrast, noisy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/29ca23d7a0397e9beaa72e3b63d17551.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-biomechanical
+ NAME_ZH: 未来生物力学
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: biomechanical style {prompt} . blend of organic and mechanical elements, futuristic, cybernetic, detailed, intricate
+ NEGATIVE_PROMPT: natural, rustic, primitive, organic, simplistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/5987795487a3fa74469499ae53963d5d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-biomechanical cyberpunk
+ NAME_ZH: 未来生物力学赛博朋克
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: biomechanical cyberpunk {prompt} . cybernetics, human-machine fusion, dystopian, organic meets artificial, dark, intricate, highly detailed
+ NEGATIVE_PROMPT: natural, colorful, deformed, sketch, low contrast, watercolor
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/75b56d7cff9b3f011248b3092ae3b8d2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-cybernetic
+ NAME_ZH: 未来赛博
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: cybernetic style {prompt} . futuristic, technological, cybernetic enhancements, robotics, artificial intelligence themes
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, historical, medieval
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/47df6db04290f01aeee7381aaf12ec07.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-cybernetic robot
+ NAME_ZH: 未来机器人
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: cybernetic robot {prompt} . android, AI, machine, metal, wires, tech, futuristic, highly detailed
+ NEGATIVE_PROMPT: organic, natural, human, sketch, watercolor, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/6f18121b4ba22c8ba7834f404501fa4c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-cyberpunk cityscape
+ NAME_ZH: 未来赛博朋克城市景观
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: cyberpunk cityscape {prompt} . neon lights, dark alleys, skyscrapers, futuristic, vibrant colors, high contrast, highly detailed
+ NEGATIVE_PROMPT: natural, rural, deformed, low contrast, black and white, sketch, watercolor
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/5e4107dace2c9dd0217a9d1b18ae28f6.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-futuristic
+ NAME_ZH: 未来主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: futuristic style {prompt} . sleek, modern, ultramodern, high tech, detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, vintage, antique
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8a61c0853ee4079bae96ad502a6ae840.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-retro cyberpunk
+ NAME_ZH: 未来复古赛博朋克
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: retro cyberpunk {prompt} . 80's inspired, synthwave, neon, vibrant, detailed, retro futurism
+ NEGATIVE_PROMPT: modern, desaturated, black and white, realism, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/460957bf9a07a58ddec08bf22d5b3698.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-retro futurism
+ NAME_ZH: 未来复古主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: retro-futuristic {prompt} . vintage sci-fi, 50s and 60s style, atomic age, vibrant, highly detailed
+ NEGATIVE_PROMPT: contemporary, realistic, rustic, primitive
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0434a7b23a0f936db615d7a9b8e805f1.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-sci-fi
+ NAME_ZH: 科幻未来主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: sci-fi style {prompt} . futuristic, technological, alien worlds, space themes, advanced civilizations
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, historical, medieval
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/7c3bde651f426273758b68a2805d0c8a.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: futuristic-vaporwave
+ NAME_ZH: 未来波
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: vaporwave style {prompt} . retro aesthetic, cyberpunk, vibrant, neon colors, vintage 80s and 90s style, highly detailed
+ NEGATIVE_PROMPT: monochrome, muted colors, realism, rustic, minimalist, dark
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/1ebb0bf67ee3ef76288b4c0a2a65d557.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-bubble bobble
+ NAME_ZH: 游戏-泡泡龙
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: Bubble Bobble style {prompt} . 8-bit, cute, pixelated, fantasy, vibrant, reminiscent of Bubble Bobble game
+ NEGATIVE_PROMPT: realistic, modern, photorealistic, violent, horror
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/19257352df33228555cb350963d9432a.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-cyberpunk game
+ NAME_ZH: 赛博朋克游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: cyberpunk game style {prompt} . neon, dystopian, futuristic, digital, vibrant, detailed, high contrast, reminiscent of cyberpunk genre video games
+ NEGATIVE_PROMPT: historical, natural, rustic, low detailed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/6165ef47bc859e3dd0e3399ba2264c2b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-fighting game
+ NAME_ZH: 格斗游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: fighting game style {prompt} . dynamic, vibrant, action-packed, detailed character design, reminiscent of fighting video games
+ NEGATIVE_PROMPT: peaceful, calm, minimalist, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/aebda69ec1097dc61966995956ae34cc.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-gta
+ NAME_ZH: 侠盗猎车手游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: GTA-style artwork {prompt} . satirical, exaggerated, pop art style, vibrant colors, iconic characters, action-packed
+ NEGATIVE_PROMPT: realistic, black and white, low contrast, impressionist, cubist, noisy, blurry, deformed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/edff7ba5fc983468e79e64beb838b829.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-mario
+ NAME_ZH: 马里奥游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: Super Mario style {prompt} . vibrant, cute, cartoony, fantasy, playful, reminiscent of Super Mario series
+ NEGATIVE_PROMPT: realistic, modern, horror, dystopian, violent
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2a229f623ce89160d1f67620599dfc7c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-minecraft
+ NAME_ZH: 我的世界游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: Minecraft style {prompt} . blocky, pixelated, vibrant colors, recognizable characters and objects, game assets
+ NEGATIVE_PROMPT: smooth, realistic, detailed, photorealistic, noise, blurry, deformed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/83e77332a9f234fd8c9cddd29e2ec3c7.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-pokemon
+ NAME_ZH: 宝可梦游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: Pokémon style {prompt} . vibrant, cute, anime, fantasy, reminiscent of Pokémon series
+ NEGATIVE_PROMPT: realistic, modern, horror, dystopian, violent
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9bfe2805e6578875a8dfe2a0ea96c281.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-retro arcade
+ NAME_ZH: 复古街机
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: retro arcade style {prompt} . 8-bit, pixelated, vibrant, classic video game, old school gaming, reminiscent of 80s and 90s arcade games
+ NEGATIVE_PROMPT: modern, ultra-high resolution, photorealistic, 3D
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/76885e0b110f1fdf3eaee52fbaf27abf.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-retro game
+ NAME_ZH: 复古游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: retro game art {prompt} . 16-bit, vibrant colors, pixelated, nostalgic, charming, fun
+ NEGATIVE_PROMPT: realistic, photorealistic, 35mm film, deformed, glitch, low contrast, noisy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/61fbae41c50f56e24dcf20fe3612d456.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-rpg fantasy game
+ NAME_ZH: 角色扮演幻想游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: role-playing game (RPG) style fantasy {prompt} . detailed, vibrant, immersive, reminiscent of high fantasy RPG games
+ NEGATIVE_PROMPT: sci-fi, modern, urban, futuristic, low detailed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/34797baf58e32b4e1a37753fe9dcff5c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-strategy game
+ NAME_ZH: 策略游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: strategy game style {prompt} . overhead view, detailed map, units, reminiscent of real-time strategy video games
+ NEGATIVE_PROMPT: first-person view, modern, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0d276527d22510a1b5f8f74eac2790df.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-streetfighter
+ NAME_ZH: 街头霸王游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: Street Fighter style {prompt} . vibrant, dynamic, arcade, 2D fighting game, highly detailed, reminiscent of Street Fighter series
+ NEGATIVE_PROMPT: 3D, realistic, modern, photorealistic, turn-based strategy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/f5927fc2b8d4a8212242bd97adcbdfa1.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: game-zelda
+ NAME_ZH: 塞尔达传说游戏
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: Legend of Zelda style {prompt} . vibrant, fantasy, detailed, epic, heroic, reminiscent of The Legend of Zelda series
+ NEGATIVE_PROMPT: sci-fi, modern, realistic, horror
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/4f9eadddbc196268258089b8cc4e5c5c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-architectural
+ NAME_ZH: 建筑
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: architectural style {prompt} . clean lines, geometric shapes, minimalist, modern, architectural drawing, highly detailed
+ NEGATIVE_PROMPT: curved lines, ornate, baroque, abstract, grunge
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/488545e9fc9417d62961f42d578547d2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-disco
+ NAME_ZH: 迪斯科
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: disco-themed {prompt} . vibrant, groovy, retro 70s style, shiny disco balls, neon lights, dance floor, highly detailed
+ NEGATIVE_PROMPT: minimalist, rustic, monochrome, contemporary, simplistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/784b0ac35e0c0fdb2df21a95d5ca1c55.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-dreamscape
+ NAME_ZH: 梦境
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: dreamscape {prompt} . surreal, ethereal, dreamy, mysterious, fantasy, highly detailed
+ NEGATIVE_PROMPT: realistic, concrete, ordinary, mundane
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2bf77804b5bf352c4e97475b1e8eb29e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-dystopian
+ NAME_ZH: 反乌托邦
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: dystopian style {prompt} . bleak, post-apocalyptic, somber, dramatic, highly detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, cheerful, optimistic, vibrant, colorful
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/065db0385e126192cbd27e22ca4f154c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-fairy tale
+ NAME_ZH: 童话故事
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: fairy tale {prompt} . magical, fantastical, enchanting, storybook style, highly detailed
+ NEGATIVE_PROMPT: realistic, modern, ordinary, mundane
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/578dfe65d9f03e83c95820d4c45cd396.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-gothic
+ NAME_ZH: 哥特
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: gothic style {prompt} . dark, mysterious, haunting, dramatic, ornate, detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, cheerful, optimistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/332a1f8cd655b724cb48bc5c91b764fc.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-grunge
+ NAME_ZH: 垃圾摇滚
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: grunge style {prompt} . textured, distressed, vintage, edgy, punk rock vibe, dirty, noisy
+ NEGATIVE_PROMPT: smooth, clean, minimalist, sleek, modern, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0ccb4c03c6d983b9f6d05ba383f9d690.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-horror
+ NAME_ZH: 恐怖
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: horror-themed {prompt} . eerie, unsettling, dark, spooky, suspenseful, grim, highly detailed
+ NEGATIVE_PROMPT: cheerful, bright, vibrant, light-hearted, cute
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/fdca4780d099a39dc50e4553622fba72.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-kawaii
+ NAME_ZH: 卡哇伊
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: kawaii style {prompt} . cute, adorable, brightly colored, cheerful, anime influence, highly detailed
+ NEGATIVE_PROMPT: dark, scary, realistic, monochrome, abstract
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/99dab4185e7337189d8959878bfa308b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-lovecraftian
+ NAME_ZH: 克苏鲁神话
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: lovecraftian horror {prompt} . eldritch, cosmic horror, unknown, mysterious, surreal, highly detailed
+ NEGATIVE_PROMPT: light-hearted, mundane, familiar, simplistic, realistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/b7c9d084ed62c1c2e1ea080ad95af61e.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-macabre
+ NAME_ZH: 恐怖的
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: macabre style {prompt} . dark, gothic, grim, haunting, highly detailed
+ NEGATIVE_PROMPT: bright, cheerful, light-hearted, cartoonish, cute
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/fd7cf2315b27b0434b909020414e1173.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-manga
+ NAME_ZH: 漫画
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: manga style {prompt} . vibrant, high-energy, detailed, iconic, Japanese comic style
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, Western comic style
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/531aac90b321d39221ba7cf70a97b232.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-metropolis
+ NAME_ZH: 大都市
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: metropolis-themed {prompt} . urban, cityscape, skyscrapers, modern, futuristic, highly detailed
+ NEGATIVE_PROMPT: rural, natural, rustic, historical, simple
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/eaad79fe3d7160cec6e0656fc3b02d3c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-minimalist
+ NAME_ZH: 极简主义
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: minimalist style {prompt} . simple, clean, uncluttered, modern, elegant
+ NEGATIVE_PROMPT: ornate, complicated, highly detailed, cluttered, disordered, messy, noisy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/e85dd1ccc7f22bdac054d74ba58521b2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-monochrome
+ NAME_ZH: 单色
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: monochrome {prompt} . black and white, contrast, tone, texture, detailed
+ NEGATIVE_PROMPT: colorful, vibrant, noisy, blurry, deformed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/cc793ada6b47561f4b63aa66dcd282ac.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-nautical
+ NAME_ZH: 航海
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: nautical-themed {prompt} . sea, ocean, ships, maritime, beach, marine life, highly detailed
+ NEGATIVE_PROMPT: landlocked, desert, mountains, urban, rustic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/3d9f6528180670804a2c5e85c7665c42.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-space
+ NAME_ZH: 太空
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: space-themed {prompt} . cosmic, celestial, stars, galaxies, nebulas, planets, science fiction, highly detailed
+ NEGATIVE_PROMPT: earthly, mundane, ground-based, realism
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/6f5742d23b43fc99cc01c8b68518df7c.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-stained glass
+ NAME_ZH: 彩色玻璃
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: stained glass style {prompt} . vibrant, beautiful, translucent, intricate, detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/ab0512ced9ed572075c49ad4f9f9a45d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-techwear fashion
+ NAME_ZH: 科技服饰
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: techwear fashion {prompt} . futuristic, cyberpunk, urban, tactical, sleek, dark, highly detailed
+ NEGATIVE_PROMPT: vintage, rural, colorful, low contrast, realism, sketch, watercolor
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/652a6bdb1b36860041222a458c7915d0.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-tribal
+ NAME_ZH: 部落
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: tribal style {prompt} . indigenous, ethnic, traditional patterns, bold, natural colors, highly detailed
+ NEGATIVE_PROMPT: modern, futuristic, minimalist, pastel
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/1ecdaf153ece5e16250b87020b7a209b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: misc-zentangle
+ NAME_ZH: 禅绕画
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: zentangle {prompt} . intricate, abstract, monochrome, patterns, meditative, highly detailed
+ NEGATIVE_PROMPT: colorful, representative, simplistic, large fields of color
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/80f3488cdb69f0ff13d83a4feac4499d.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-collage
+ NAME_ZH: 纸艺拼贴
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: collage style {prompt} . mixed media, layered, textural, detailed, artistic
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/c9bbdaba28358faf1a8c1df44b3de950.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-flat papercut
+ NAME_ZH: 平面剪纸
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: flat papercut style {prompt} . silhouette, clean cuts, paper, sharp edges, minimalist, color block
+ NEGATIVE_PROMPT: 3D, high detail, noise, grainy, blurry, painting, drawing, photo, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/1bffcd9a4086b29c367b8301f26eed3b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-kirigami
+ NAME_ZH: 剪纸
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: kirigami representation of {prompt} . 3D, paper folding, paper cutting, Japanese, intricate, symmetrical, precision, clean lines
+ NEGATIVE_PROMPT: painting, drawing, 2D, noisy, blurry, deformed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2e5f95977b97b2a6248924e118a23076.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-paper mache
+ NAME_ZH: 纸浆塑型
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: paper mache representation of {prompt} . 3D, sculptural, textured, handmade, vibrant, fun
+ NEGATIVE_PROMPT: 2D, flat, photo, sketch, digital art, deformed, noisy, blurry
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/68441f668fcf7a51db5fc9b9d4d45ce7.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-paper quilling
+ NAME_ZH: 纸卷艺术
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: paper quilling art of {prompt} . intricate, delicate, curling, rolling, shaping, coiling, loops, 3D, dimensional, ornamental
+ NEGATIVE_PROMPT: photo, painting, drawing, 2D, flat, deformed, noisy, blurry
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8c0504eb14341b66796612c46bb3748b.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-papercut collage
+ NAME_ZH: 剪纸拼贴
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: papercut collage of {prompt} . mixed media, textured paper, overlapping, asymmetrical, abstract, vibrant
+ NEGATIVE_PROMPT: photo, 3D, realistic, drawing, painting, high detail, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/217e7c8b905b63c63c94bd6ff8fe3477.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-papercut shadow box
+ NAME_ZH: 剪纸影箱
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: 3D papercut shadow box of {prompt} . layered, dimensional, depth, silhouette, shadow, papercut, handmade, high contrast
+ NEGATIVE_PROMPT: painting, drawing, photo, 2D, flat, high detail, blurry, noisy, disfigured
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2af0d61de1b17a5a5257f512549f9616.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-stacked papercut
+ NAME_ZH: 堆叠剪纸
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: stacked papercut art of {prompt} . 3D, layered, dimensional, depth, precision cut, stacked layers, papercut, high contrast
+ NEGATIVE_PROMPT: 2D, flat, noisy, blurry, painting, drawing, photo, deformed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/62618e32e90539827e0eb279a6ea170a.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: papercraft-thick layered papercut
+ NAME_ZH: 厚层剪纸
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: thick layered papercut art of {prompt} . deep 3D, volumetric, dimensional, depth, thick paper, high stack, heavy texture, tangible layers
+ NEGATIVE_PROMPT: 2D, flat, thin paper, low stack, smooth texture, painting, drawing, photo, deformed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/17008c40a6be8d41d3921ee88bac27ad.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-alien
+ NAME_ZH: 异形
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: alien-themed {prompt} . extraterrestrial, cosmic, otherworldly, mysterious, sci-fi, highly detailed
+ NEGATIVE_PROMPT: earthly, mundane, common, realistic, simple
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/eae6043dee2b2d94cd9a22cb044db8b3.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-film noir
+ NAME_ZH: 黑色电影
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: film noir style {prompt} . monochrome, high contrast, dramatic shadows, 1940s style, mysterious, cinematic
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, vibrant, colorful
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0b3629be5ebb7cda463877a1993a0298.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-glamour
+ NAME_ZH: 魅力
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: glamorous photo {prompt} . high fashion, luxurious, extravagant, stylish, sensual, opulent, elegance, stunning beauty, professional, high contrast, detailed
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, distorted, grainy, sketch, low contrast, dull, plain, modest
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/f85dc6aa5d7f7d6d3663f2270e276796.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-hdr
+ NAME_ZH: 高动态范围
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: HDR photo of {prompt} . High dynamic range, vivid, rich details, clear shadows and highlights, realistic, intense, enhanced contrast, highly detailed
+ NEGATIVE_PROMPT: flat, low contrast, oversaturated, underexposed, overexposed, blurred, noisy
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/269a98b284ae1345b0e5ac1f85c33179.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-iphone photographic
+ NAME_ZH: iPhone摄影
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: iphone photo {prompt} . large depth of field, deep depth of field, highly detailed
+ NEGATIVE_PROMPT: drawing, painting, crayon, sketch, graphite, impressionist, noisy, blurry, soft, deformed, ugly, shallow depth of field, bokeh
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/48c556a6b3c3a533847b80b031cd4962.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-long exposure
+ NAME_ZH: 长曝光
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: long exposure photo of {prompt} . Blurred motion, streaks of light, surreal, dreamy, ghosting effect, highly detailed
+ NEGATIVE_PROMPT: static, noisy, deformed, shaky, abrupt, flat, low contrast
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/fda88c3fc43ea8cc04ce9b5c0261627f.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-neon noir
+ NAME_ZH: 霓虹黑色
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: neon noir {prompt} . cyberpunk, dark, rainy streets, neon signs, high contrast, low light, vibrant, highly detailed
+ NEGATIVE_PROMPT: bright, sunny, daytime, low contrast, black and white, sketch, watercolor
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8a5b71434592771e064af4b9017b59d6.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-silhouette
+ NAME_ZH: 剪影
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: silhouette style {prompt} . high contrast, minimalistic, black and white, stark, dramatic
+ NEGATIVE_PROMPT: ugly, deformed, noisy, blurry, low contrast, color, realism, photorealistic
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/1892a37e66d33cdd9b713aaab6b8c597.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: photo-tilt-shift
+ NAME_ZH: 倾斜移位
+ DESCRIPTION:
+ SOURCE: twri
+ PROMPT: tilt-shift photo of {prompt} . selective focus, miniature effect, blurred background, highly detailed, vibrant, perspective control
+ NEGATIVE_PROMPT: blurry, noisy, deformed, flat, low contrast, unrealistic, oversaturated, underexposed
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/d058a157de848a5c03d9a3e1e0e560a2.jpg
+ PROMPT_EXAMPLE: a boy wearing green jacket
\ No newline at end of file
diff --git a/scepter/workflow/config/pixart_aplha_pro.yaml b/scepter/workflow/config/pixart_aplha_pro.yaml
new file mode 100644
index 0000000..9175bd1
--- /dev/null
+++ b/scepter/workflow/config/pixart_aplha_pro.yaml
@@ -0,0 +1,266 @@
+NAME: PIXART_ALPHA
+IS_DEFAULT: False
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[1024, 1024]]
+ INPUT:
+ IMAGE:
+ ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
+ TARGET_SIZE_AS_TUPLE: [1024, 1024]
+ PROMPT: ""
+ NEGATIVE_PROMPT: ""
+ PROMPT_PREFIX: ""
+ SAMPLE: ddim
+ SAMPLE_STEPS: 20
+ GUIDE_SCALE: 4.5
+ GUIDE_RESCALE: 0.5
+ DISCRETIZATION: trailing
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float32
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: float32
+ INPUT: ["LATENT"]
+ PARAS:
+ SCALE_FACTOR: 0.18215
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: float16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float32
+ INPUT: ["PROMPT"]
+#
+MODEL:
+ PRETRAINED_MODEL:
+ DECODER_BIAS: 0.5
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ "NAME": "linear"
+ "BETA_MIN": 0.0001
+ "BETA_MAX": 0.02
+ #
+ DIFFUSION_MODEL:
+ NAME: PixArt
+ PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@PixArt-XL-2-1024-MS.pth
+ INPUT_SIZE: 128
+ PATCH_SIZE: 2
+ IN_CHANNELS: 4
+ HIDDEN_SIZE: 1152
+ DEPTH: 28
+ NUM_HEADS: 16
+ MLP_RATIO: 4.0
+ CLASS_DROPOUT_PROB: 0.1
+ PRED_SIGMA: True
+ DROP_PATH: 0.0
+ WINDOW_DIZE: 0
+ USE_REL_POS: False
+ CAPTION_CHANNELS: 4096
+ LEWEI_SCALE: 2
+ MODEL_MAX_LENGTH: 120
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-base@512-base-ema.safetensors
+ EMBED_DIM: 4
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 1
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5EmbedderHF
+ PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
+ TOKENIZER_PATH: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
+ LENGTH: 120
+ CLEAN: heavy
+ USE_GRAD: False
+#
+MODEL_LOCAL:
+ PRETRAINED_MODEL:
+ DECODER_BIAS: 0.5
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ "NAME": "linear"
+ "BETA_MIN": 0.0001
+ "BETA_MAX": 0.02
+ #
+ DIFFUSION_MODEL:
+ NAME: PixArt
+ PRETRAINED_MODEL: models/scepter/PixArt-alpha/PixArt-XL-2-1024-MS.pth
+ INPUT_SIZE: 128
+ PATCH_SIZE: 2
+ IN_CHANNELS: 4
+ HIDDEN_SIZE: 1152
+ DEPTH: 28
+ NUM_HEADS: 16
+ MLP_RATIO: 4.0
+ CLASS_DROPOUT_PROB: 0.1
+ PRED_SIGMA: True
+ DROP_PATH: 0.0
+ WINDOW_DIZE: 0
+ USE_REL_POS: False
+ CAPTION_CHANNELS: 4096
+ LEWEI_SCALE: 2
+ MODEL_MAX_LENGTH: 120
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-2-base/512-base-ema.safetensors
+ EMBED_DIM: 4
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 1
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5EmbedderHF
+ PRETRAINED_MODEL: models/scepter/PixArt-alpha/t5-v1_1-xxl/
+ TOKENIZER_PATH: models/scepter/PixArt-alpha/t5-v1_1-xxl/
+ LENGTH: 120
+ CLEAN: heavy
+ USE_GRAD: False
+#
+MODEL_HF:
+ PRETRAINED_MODEL:
+ DECODER_BIAS: 0.5
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ "NAME": "linear"
+ "BETA_MIN": 0.0001
+ "BETA_MAX": 0.02
+ #
+ DIFFUSION_MODEL:
+ NAME: PixArt
+ PRETRAINED_MODEL: hf://PixArt-alpha/PixArt-alpha@PixArt-XL-2-1024-MS.pth
+ INPUT_SIZE: 128
+ PATCH_SIZE: 2
+ IN_CHANNELS: 4
+ HIDDEN_SIZE: 1152
+ DEPTH: 28
+ NUM_HEADS: 16
+ MLP_RATIO: 4.0
+ CLASS_DROPOUT_PROB: 0.1
+ PRED_SIGMA: True
+ DROP_PATH: 0.0
+ WINDOW_DIZE: 0
+ USE_REL_POS: False
+ CAPTION_CHANNELS: 4096
+ LEWEI_SCALE: 2
+ MODEL_MAX_LENGTH: 120
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ PRETRAINED_MODEL: hf://stabilityai/stable-diffusion-2-base@512-base-ema.safetensors
+ EMBED_DIM: 4
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 1
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: T5EmbedderHF
+ PRETRAINED_MODEL: hf://PixArt-alpha/PixArt-alpha@t5-v1_1-xxl/
+ TOKENIZER_PATH: hf://PixArt-alpha/PixArt-alpha@t5-v1_1-xxl/
+ LENGTH: 120
+ CLEAN: heavy
+ USE_GRAD: False
\ No newline at end of file
diff --git a/scepter/workflow/config/scepter_workflow.yaml b/scepter/workflow/config/scepter_workflow.yaml
new file mode 100644
index 0000000..f617c9f
--- /dev/null
+++ b/scepter/workflow/config/scepter_workflow.yaml
@@ -0,0 +1,127 @@
+FILE_SYSTEMS:
+ -
+ NAME: "ModelscopeFs"
+ TEMP_DIR: "models/scepter"
+ ENABLE_MD5_PATH: False
+ -
+ NAME: "HttpFs"
+ TEMP_DIR: "models/scepter"
+ ENABLE_MD5_PATH: False
+ -
+ NAME: "HuggingfaceFs"
+ TEMP_DIR: "models/scepter"
+ ENABLE_MD5_PATH: False
+ -
+ NAME: "LocalFs"
+ TEMP_DIR: "models/scepter"
+ ENABLE_MD5_PATH: False
+
+BASE_MODELS:
+ -
+ NAME: SD_XL1.0
+ DIFFUSION_MODEL: SD_XL1.0_DiffusionUNetXL
+ FIRST_STAGE_MODEL: SD_XL1.0_AutoencoderKL
+ COND_STAGE_MODEL: SD_XL1.0_GeneralConditioner
+ CONFIG: config/sdxl1.0_pro.yaml
+ -
+ NAME: SD1.5
+ DIFFUSION_MODEL: SD1.5_DiffusionUNet
+ FIRST_STAGE_MODEL: SD1.5_AutoencoderKL
+ COND_STAGE_MODEL: SD1.5_FrozenCLIPEmbedder
+ CONFIG: config/sd15_pro.yaml
+ -
+ NAME: PIXART
+ DIFFUSION_MODEL: PIXART_ALPHA_PixArt
+ FIRST_STAGE_MODEL: PIXART_ALPHA_AutoencoderKL
+ COND_STAGE_MODEL: PIXART_ALPHA_T5EmbedderHF
+ CONFIG: config/pixart_aplha_pro.yaml
+ -
+ NAME: SD3
+ DIFFUSION_MODEL: SD3_MMDiT
+ FIRST_STAGE_MODEL: SD3_AutoencoderKL
+ COND_STAGE_MODEL: SD3_SD3TextEmbedder
+ CONFIG: config/sd3_pro.yaml
+ -
+ NAME: FLUX1.0_DEV
+ DIFFUSION_MODEL: FLUX1.0_DEV_Flux
+ FIRST_STAGE_MODEL: FLUX1.0_DEV_AutoencoderKLFlux
+ COND_STAGE_MODEL: FLUX1.0_DEV_T5PlusClipFluxEmbedder
+ CONFIG: config/flux1.0_dev_pro.yaml
+ -
+ NAME: FLUX1.0_SCHNELL
+ DIFFUSION_MODEL: FLUX1.0_SCHNELL_Flux
+ FIRST_STAGE_MODEL: FLUX1.0_SCHNELL_AutoencoderKLFlux
+ COND_STAGE_MODEL: FLUX1.0_SCHNELL_T5PlusClipFluxEmbedder
+ CONFIG: config/flux1.0_schnell_pro.yaml
+
+MODEL_SOURCE:
+ - "ModelScope"
+ - "HuggingFace"
+ - "Local"
+
+BASE_PARAMETERS:
+ SAMPLER:
+ - "ddim"
+ - "euler"
+ - "euler_ancestral"
+ - "henu"
+ - "dpm2"
+ - "dpm2_ancestral"
+ - "dpmpp_2m"
+ - "dpmpp_sde"
+ - "dpmpp_2m_sde"
+ - "dpmpp_2s_ancestral"
+ - "dpm2_karras"
+ - "dpm2_ancestral_karras"
+ - "dpmpp_2s_ancestral_karras"
+ - "dpmpp_2m_karras"
+ - "dpmpp_sde_karras"
+ - "dpmpp_2m_sde_karras"
+ - "flow_eluer"
+
+ DISCRETIZATION:
+ - "trailing"
+ - "leading"
+ - "linspace"
+
+ OUTPUT_HEIGHT:
+ - 1024
+ - 512
+ - 704
+ - 720
+ - 768
+ - 832
+ - 896
+ - 960
+ - 1088
+ - 1152
+ - 1216
+ - 1280
+ - 1344
+ - 1408
+ - 1472
+ - 1536
+ - 1600
+ - 1664
+ - 1728
+
+ OUTPUT_WIDTH:
+ - 1024
+ - 512
+ - 704
+ - 720
+ - 768
+ - 832
+ - 896
+ - 960
+ - 1088
+ - 1152
+ - 1216
+ - 1280
+ - 1344
+ - 1408
+ - 1472
+ - 1536
+ - 1600
+ - 1664
+ - 1728
\ No newline at end of file
diff --git a/scepter/workflow/config/sd15_pro.yaml b/scepter/workflow/config/sd15_pro.yaml
new file mode 100644
index 0000000..445c909
--- /dev/null
+++ b/scepter/workflow/config/sd15_pro.yaml
@@ -0,0 +1,281 @@
+NAME: SD1.5
+IS_DEFAULT: False
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[512, 512]]
+ INPUT:
+ IMAGE:
+ PROMPT: ""
+ NEGATIVE_PROMPT: ""
+ TARGET_SIZE_AS_TUPLE: [512, 512]
+ PROMPT_PREFIX: ""
+ SAMPLE: ddim
+ SAMPLE_STEPS: 50
+ GUIDE_SCALE: 7.5
+ GUIDE_RESCALE: 0.5
+ DISCRETIZATION: trailing
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float16
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: float16
+ INPUT: ["LATENT"]
+ PARAS:
+ # SCALE_FACTOR DESCRIPTION: The vae embeding scale. TYPE: float default: 0.18215
+ SCALE_FACTOR: 0.18215
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: float16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode_text
+ DTYPE: float16
+ INPUT: ["PROMPT", "NEGATIVE_PROMPT"]
+#
+MODEL:
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-v1-5@v1-5-pruned-emaonly.safetensors
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ NAME: "scaled_linear"
+ BETA_MIN: 0.00085
+ BETA_MAX: 0.0120
+ #
+ DIFFUSION_MODEL:
+ NAME: DiffusionUNet
+ IN_CHANNELS: 4
+ OUT_CHANNELS: 4
+ MODEL_CHANNELS: 320
+ NUM_HEADS: 8
+ NUM_RES_BLOCKS: 2
+ ATTENTION_RESOLUTIONS: [ 4, 2, 1 ]
+ CHANNEL_MULT: [ 1, 2, 4, 4 ]
+ CONV_RESAMPLE: True
+ DIMS: 2
+ USE_CHECKPOINT: False
+ USE_SCALE_SHIFT_NORM: False
+ RESBLOCK_UPDOWN: False
+ USE_SPATIAL_TRANSFORMER: True
+ TRANSFORMER_DEPTH: 1
+ CONTEXT_DIM: 768
+ DISABLE_MIDDLE_SELF_ATTN: False
+ USE_LINEAR_IN_TRANSFORMER: False
+ IGNORE_KEYS: []
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ EMBED_DIM: 4
+ IGNORE_KEYS: []
+ BATCH_SIZE: 4
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ TOKENIZER:
+ NAME: ClipTokenizer
+ PRETRAINED_PATH: ms://AI-ModelScope/clip-vit-large-patch14
+ LENGTH: 77
+ CLEAN: True
+ #
+ COND_STAGE_MODEL:
+ NAME: FrozenCLIPEmbedder
+ FREEZE: True
+ USE_GRAD: False
+ LAYER: last
+ PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
+#
+MODEL_LOCAL:
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-v1-5/v1-5-pruned-emaonly.safetensors
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ NAME: "scaled_linear"
+ BETA_MIN: 0.00085
+ BETA_MAX: 0.0120
+ #
+ DIFFUSION_MODEL:
+ NAME: DiffusionUNet
+ IN_CHANNELS: 4
+ OUT_CHANNELS: 4
+ MODEL_CHANNELS: 320
+ NUM_HEADS: 8
+ NUM_RES_BLOCKS: 2
+ ATTENTION_RESOLUTIONS: [ 4, 2, 1 ]
+ CHANNEL_MULT: [ 1, 2, 4, 4 ]
+ CONV_RESAMPLE: True
+ DIMS: 2
+ USE_CHECKPOINT: False
+ USE_SCALE_SHIFT_NORM: False
+ RESBLOCK_UPDOWN: False
+ USE_SPATIAL_TRANSFORMER: True
+ TRANSFORMER_DEPTH: 1
+ CONTEXT_DIM: 768
+ DISABLE_MIDDLE_SELF_ATTN: False
+ USE_LINEAR_IN_TRANSFORMER: False
+ IGNORE_KEYS: []
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ EMBED_DIM: 4
+ IGNORE_KEYS: []
+ BATCH_SIZE: 4
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ TOKENIZER:
+ NAME: ClipTokenizer
+ PRETRAINED_PATH: models/scepter/clip-vit-large-patch14
+ LENGTH: 77
+ CLEAN: True
+ #
+ COND_STAGE_MODEL:
+ NAME: FrozenCLIPEmbedder
+ FREEZE: True
+ USE_GRAD: False
+ LAYER: last
+ PRETRAINED_MODEL: models/scepter/clip-vit-large-patch14
+#
+MODEL_HF:
+ PRETRAINED_MODEL: hf://stable-diffusion-v1-5/stable-diffusion-v1-5@v1-5-pruned-emaonly.safetensors
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ NAME: "scaled_linear"
+ BETA_MIN: 0.00085
+ BETA_MAX: 0.0120
+ #
+ DIFFUSION_MODEL:
+ NAME: DiffusionUNet
+ IN_CHANNELS: 4
+ OUT_CHANNELS: 4
+ MODEL_CHANNELS: 320
+ NUM_HEADS: 8
+ NUM_RES_BLOCKS: 2
+ ATTENTION_RESOLUTIONS: [ 4, 2, 1 ]
+ CHANNEL_MULT: [ 1, 2, 4, 4 ]
+ CONV_RESAMPLE: True
+ DIMS: 2
+ USE_CHECKPOINT: False
+ USE_SCALE_SHIFT_NORM: False
+ RESBLOCK_UPDOWN: False
+ USE_SPATIAL_TRANSFORMER: True
+ TRANSFORMER_DEPTH: 1
+ CONTEXT_DIM: 768
+ DISABLE_MIDDLE_SELF_ATTN: False
+ USE_LINEAR_IN_TRANSFORMER: False
+ IGNORE_KEYS: []
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ EMBED_DIM: 4
+ IGNORE_KEYS: []
+ BATCH_SIZE: 4
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ TOKENIZER:
+ NAME: ClipTokenizer
+ PRETRAINED_PATH: hf://openai/clip-vit-large-patch14
+ LENGTH: 77
+ CLEAN: True
+ #
+ COND_STAGE_MODEL:
+ NAME: FrozenCLIPEmbedder
+ FREEZE: True
+ USE_GRAD: False
+ LAYER: last
+ PRETRAINED_MODEL: hf://openai/clip-vit-large-patch14
diff --git a/scepter/workflow/config/sd3_pro.yaml b/scepter/workflow/config/sd3_pro.yaml
new file mode 100644
index 0000000..09f07ab
--- /dev/null
+++ b/scepter/workflow/config/sd3_pro.yaml
@@ -0,0 +1,348 @@
+NAME: SD3
+IS_DEFAULT: False
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[1024, 1024]]
+ INPUT:
+ IMAGE:
+ ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
+ TARGET_SIZE_AS_TUPLE: [1024, 1024]
+ PROMPT: ""
+ NEGATIVE_PROMPT: ""
+ PROMPT_PREFIX: ""
+ SAMPLE: euler
+ SAMPLE_STEPS: 28
+ GUIDE_SCALE: 5.0
+ GUIDE_RESCALE: 0.0
+ DISCRETIZATION: trailing
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float32
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: float32
+ INPUT: ["LATENT"]
+ PARAS:
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: float16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float32
+ INPUT: ["PROMPT"]
+#
+MODEL:
+ PRETRAINED_MODEL:
+ SCHEDULE:
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ DEFAULT_N_PROMPT:
+ SCHEDULE_ARGS:
+ "NAME": "shifted"
+ "SHIFT": 3
+ T_WEIGHT: uniform
+ #
+ DIFFUSION_MODEL:
+ NAME: MMDiT
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
+ IGNORE_KEYS: '^first_stage_model.'
+ IN_CHANNELS: 16
+ PATCH_SIZE: 2
+ OUT_CHANNELS: 16
+ DEPTH: 24
+ INPUT_SIZE:
+ ADM_IN_CHANNELS: 2048
+ CONTEXT_EMBEDDER_CONFIG: { 'target': 'torch.nn.Linear', 'params': { 'in_features': 4096, 'out_features': 1536 } }
+ NUM_PATCHES: 36864
+ POS_EMBED_MAX_SIZE: 192
+ POS_EMBED_SCALING_FACTOR:
+ USE_CHECKPOINT: True
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
+ EMBED_DIM: 16
+ IGNORE_KEYS: '^model.diffusion_model.'
+ BATCH_SIZE: 1
+ USE_CONV: False
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: SD3TextEmbedder
+ P_ZERO: 0.0
+ CLIP_L:
+ NAME: FrozenCLIPEmbedder2
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder
+ TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: penultimate
+ RETURN_POOLED: True
+ USE_FINAL_LAYER_NORM: False
+ IS_TRAINABLE: False
+ CLIP_G:
+ NAME: FrozenCLIPEmbedder2
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_2
+ TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_2
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: penultimate
+ RETURN_POOLED: True
+ USE_FINAL_LAYER_NORM: False
+ IS_TRAINABLE: False
+ T5_XXL:
+ NAME: T5EmbedderHF
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_3
+ TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_3
+ LENGTH: 256
+ CLEAN: whitespace
+ USE_GRAD: False
+ T5_DTYPE: float16
+#
+MODEL_LOCAL:
+ PRETRAINED_MODEL:
+ SCHEDULE:
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ DEFAULT_N_PROMPT:
+ SCHEDULE_ARGS:
+ "NAME": "shifted"
+ "SHIFT": 3
+ T_WEIGHT: uniform
+ #
+ DIFFUSION_MODEL:
+ NAME: MMDiT
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-3-medium/sd3_medium.safetensors
+ IGNORE_KEYS: '^first_stage_model.'
+ IN_CHANNELS: 16
+ PATCH_SIZE: 2
+ OUT_CHANNELS: 16
+ DEPTH: 24
+ INPUT_SIZE:
+ ADM_IN_CHANNELS: 2048
+ CONTEXT_EMBEDDER_CONFIG: { 'target': 'torch.nn.Linear', 'params': { 'in_features': 4096, 'out_features': 1536 } }
+ NUM_PATCHES: 36864
+ POS_EMBED_MAX_SIZE: 192
+ POS_EMBED_SCALING_FACTOR:
+ USE_CHECKPOINT: True
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-3-medium/sd3_medium.safetensors
+ EMBED_DIM: 16
+ IGNORE_KEYS: '^model.diffusion_model.'
+ BATCH_SIZE: 1
+ USE_CONV: False
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: SD3TextEmbedder
+ P_ZERO: 0.0
+ CLIP_L:
+ NAME: FrozenCLIPEmbedder2
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-3-medium-diffusers/text_encoder
+ TOKENIZER_PATH: models/scepter/stable-diffusion-3-medium-diffusers/tokenizer
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: penultimate
+ RETURN_POOLED: True
+ USE_FINAL_LAYER_NORM: False
+ IS_TRAINABLE: False
+ CLIP_G:
+ NAME: FrozenCLIPEmbedder2
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-3-medium-diffusers/text_encoder_2
+ TOKENIZER_PATH: models/scepter/stable-diffusion-3-medium-diffusers/tokenizer_2
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: penultimate
+ RETURN_POOLED: True
+ USE_FINAL_LAYER_NORM: False
+ IS_TRAINABLE: False
+ T5_XXL:
+ NAME: T5EmbedderHF
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-3-medium-diffusers/text_encoder_3
+ TOKENIZER_PATH: models/scepter/stable-diffusion-3-medium-diffusers/tokenizer_3
+ LENGTH: 256
+ CLEAN: whitespace
+ USE_GRAD: False
+ T5_DTYPE: float16
+#
+MODEL_HF:
+ PRETRAINED_MODEL:
+ SCHEDULE:
+ PARAMETERIZATION: rf
+ TIMESTEPS: 1000
+ MIN_SNR_GAMMA:
+ ZERO_TERMINAL_SNR: False
+ PRETRAINED_MODEL:
+ IGNORE_KEYS: [ ]
+ SCALE_FACTOR: 1.5305
+ SHIFT_FACTOR: 0.0609
+ DEFAULT_N_PROMPT:
+ SCHEDULE_ARGS:
+ "NAME": "shifted"
+ "SHIFT": 3
+ T_WEIGHT: uniform
+ #
+ DIFFUSION_MODEL:
+ NAME: MMDiT
+ PRETRAINED_MODEL: hf://stabilityai/stable-diffusion-3-medium@sd3_medium.safetensors
+ IGNORE_KEYS: '^first_stage_model.'
+ IN_CHANNELS: 16
+ PATCH_SIZE: 2
+ OUT_CHANNELS: 16
+ DEPTH: 24
+ INPUT_SIZE:
+ ADM_IN_CHANNELS: 2048
+ CONTEXT_EMBEDDER_CONFIG: { 'target': 'torch.nn.Linear', 'params': { 'in_features': 4096, 'out_features': 1536 } }
+ NUM_PATCHES: 36864
+ POS_EMBED_MAX_SIZE: 192
+ POS_EMBED_SCALING_FACTOR:
+ USE_CHECKPOINT: True
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ PRETRAINED_MODEL: hf://stabilityai/stable-diffusion-3-medium@sd3_medium.safetensors
+ EMBED_DIM: 16
+ IGNORE_KEYS: '^model.diffusion_model.'
+ BATCH_SIZE: 1
+ USE_CONV: False
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 16
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: SD3TextEmbedder
+ P_ZERO: 0.0
+ CLIP_L:
+ NAME: FrozenCLIPEmbedder2
+ PRETRAINED_MODEL: hf://stabilityai/stable-diffusion-3-medium-diffusers@text_encoder
+ TOKENIZER_PATH: hf://stabilityai/stable-diffusion-3-medium-diffusers@tokenizer
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: penultimate
+ RETURN_POOLED: True
+ USE_FINAL_LAYER_NORM: False
+ IS_TRAINABLE: False
+ CLIP_G:
+ NAME: FrozenCLIPEmbedder2
+ PRETRAINED_MODEL: hf://stabilityai/stable-diffusion-3-medium-diffusers@text_encoder_2
+ TOKENIZER_PATH: hf://stabilityai/stable-diffusion-3-medium-diffusers@tokenizer_2
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: penultimate
+ RETURN_POOLED: True
+ USE_FINAL_LAYER_NORM: False
+ IS_TRAINABLE: False
+ T5_XXL:
+ NAME: T5EmbedderHF
+ PRETRAINED_MODEL: hf://stabilityai/stable-diffusion-3-medium-diffusers@text_encoder_3
+ TOKENIZER_PATH: hf://stabilityai/stable-diffusion-3-medium-diffusers@tokenizer_3
+ LENGTH: 256
+ CLEAN: whitespace
+ USE_GRAD: False
+ T5_DTYPE: float16
\ No newline at end of file
diff --git a/scepter/workflow/config/sdxl1.0_pro.yaml b/scepter/workflow/config/sdxl1.0_pro.yaml
new file mode 100644
index 0000000..0b96aa3
--- /dev/null
+++ b/scepter/workflow/config/sdxl1.0_pro.yaml
@@ -0,0 +1,419 @@
+NAME: SD_XL1.0
+IS_DEFAULT: True
+DEFAULT_PARAS:
+ PARAS:
+ RESOLUTIONS: [[1024, 1024]]
+ INPUT:
+ IMAGE:
+ ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
+ TARGET_SIZE_AS_TUPLE: [1024, 1024]
+ AESTHETIC_SCORE: 6.0
+ NEGATIVE_AESTHETIC_SCORE: 2.5
+ PROMPT: ""
+ NEGATIVE_PROMPT: ""
+ PROMPT_PREFIX: ""
+ CROP_COORDS_TOP_LEFT: [0, 0]
+ SAMPLE: ddim
+ SAMPLE_STEPS: 50
+ GUIDE_SCALE: 7.5
+ GUIDE_RESCALE: 0.5
+ DISCRETIZATION: trailing
+ REFINE_SAMPLE: ddim
+ REFINE_GUIDE_SCALE: 7.5
+ REFINE_GUIDE_RESCALE: 0.5
+ REFINE_DISCRETIZATION: trailing
+ OUTPUT:
+ LATENT:
+ BEFORE_REFINE_IMAGES:
+ IMAGES:
+ SEED:
+ MODULES_PARAS:
+ FIRST_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float32
+ INPUT: ["IMAGE"]
+ -
+ NAME: decode
+ DTYPE: float32
+ INPUT: ["LATENT"]
+ PARAS:
+ # SCALE_FACTOR DESCRIPTION: The vae embeding scale. TYPE: float default: 0.18215
+ SCALE_FACTOR: 0.13025
+ SIZE_FACTOR: 8
+ DIFFUSION_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: float16
+ INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
+ COND_STAGE_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float16
+ INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"]
+ REFINER_MODEL:
+ FUNCTION:
+ -
+ NAME: forward
+ DTYPE: float16
+ INPUT: ["SAMPLE_STEPS", "REFINE_SAMPLE", "REFINE_GUIDE_SCALE", "REFINE_GUIDE_RESCALE", "REFINE_DISCRETIZATION"]
+ REFINER_COND_MODEL:
+ FUNCTION:
+ -
+ NAME: encode
+ DTYPE: float16
+ INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "AESTHETIC_SCORE", "NEGATIVE_AESTHETIC_SCORE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"]
+#
+MODEL:
+ PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-xl-base-1.0@sd_xl_base_1.0.safetensors
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ NAME: "scaled_linear"
+ BETA_MIN: 0.00085
+ BETA_MAX: 0.0120
+ DIFFUSION_MODEL:
+ NAME: DiffusionUNetXL
+ PRETRAINED_MODEL:
+ IN_CHANNELS: 4
+ OUT_CHANNELS: 4
+ NUM_RES_BLOCKS: 2
+ MODEL_CHANNELS: 320
+ ATTENTION_RESOLUTIONS: [4, 2]
+ DROPOUT: 0
+ CHANNEL_MULT: [1, 2, 4]
+ CONV_RESAMPLE: True
+ DIMS: 2
+ NUM_CLASSES: sequential
+ USE_CHECKPOINT: False
+ NUM_HEADS: -1
+ NUM_HEADS_CHANNELS: 64
+ USE_SCALE_SHIFT_NORM: False
+ RESBLOCK_UPDOWN: False
+ USE_NEW_ATTENTION_ORDER: True
+ USE_SPATIAL_TRANSFORMER: True
+ TRANSFORMER_DEPTH: [1, 2, 10]
+ CONTEXT_DIM: 2048
+ DISABLE_MIDDLE_SELF_ATTN: False
+ USE_LINEAR_IN_TRANSFORMER: True
+ ADM_IN_CHANNELS: 2816
+ USE_SENTENCE_EMB: False
+ USE_WORD_MAPPING: False
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ EMBED_DIM: 4
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 1
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: GeneralConditioner
+ USE_GRAD: False
+ EMBEDDERS:
+ -
+ NAME: FrozenCLIPEmbedder
+ PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
+ TOKENIZER_PATH: ms://AI-ModelScope/clip-vit-large-patch14
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: hidden
+ LAYER_IDX: 11
+ USE_FINAL_LAYER_NORM: False
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["prompt"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: FrozenOpenCLIPEmbedder2
+ ARCH: ViT-bigG-14
+ MAX_LENGTH: 77
+ FREEZE: True
+ ALWAYS_RETURN_POOLED: True
+ LEGACY: False
+ LAYER: penultimate
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["prompt"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["original_size_as_tuple"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["crop_coords_top_left"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["target_size_as_tuple"]
+ LEGACY_UCG_VALUE:
+#
+MODEL_LOCAL:
+ PRETRAINED_MODEL: models/scepter/stable-diffusion-xl-base-1.0/sd_xl_base_1.0.safetensors
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ NAME: "scaled_linear"
+ BETA_MIN: 0.00085
+ BETA_MAX: 0.0120
+ DIFFUSION_MODEL:
+ NAME: DiffusionUNetXL
+ PRETRAINED_MODEL:
+ IN_CHANNELS: 4
+ OUT_CHANNELS: 4
+ NUM_RES_BLOCKS: 2
+ MODEL_CHANNELS: 320
+ ATTENTION_RESOLUTIONS: [4, 2]
+ DROPOUT: 0
+ CHANNEL_MULT: [1, 2, 4]
+ CONV_RESAMPLE: True
+ DIMS: 2
+ NUM_CLASSES: sequential
+ USE_CHECKPOINT: False
+ NUM_HEADS: -1
+ NUM_HEADS_CHANNELS: 64
+ USE_SCALE_SHIFT_NORM: False
+ RESBLOCK_UPDOWN: False
+ USE_NEW_ATTENTION_ORDER: True
+ USE_SPATIAL_TRANSFORMER: True
+ TRANSFORMER_DEPTH: [1, 2, 10]
+ CONTEXT_DIM: 2048
+ DISABLE_MIDDLE_SELF_ATTN: False
+ USE_LINEAR_IN_TRANSFORMER: True
+ ADM_IN_CHANNELS: 2816
+ USE_SENTENCE_EMB: False
+ USE_WORD_MAPPING: False
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ EMBED_DIM: 4
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 1
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: GeneralConditioner
+ USE_GRAD: False
+ EMBEDDERS:
+ -
+ NAME: FrozenCLIPEmbedder
+ PRETRAINED_MODEL: models/scepter/clip-vit-large-patch14
+ TOKENIZER_PATH: models/scepter/clip-vit-large-patch14
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: hidden
+ LAYER_IDX: 11
+ USE_FINAL_LAYER_NORM: False
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["prompt"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: FrozenOpenCLIPEmbedder2
+ ARCH: ViT-bigG-14
+ MAX_LENGTH: 77
+ FREEZE: True
+ ALWAYS_RETURN_POOLED: True
+ LEGACY: False
+ LAYER: penultimate
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["prompt"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["original_size_as_tuple"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["crop_coords_top_left"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["target_size_as_tuple"]
+ LEGACY_UCG_VALUE:
+#
+MODEL_HF:
+ PRETRAINED_MODEL: hf://stabilityai/stable-diffusion-xl-base-1.0@sd_xl_base_1.0.safetensors
+ SCHEDULE:
+ PARAMETERIZATION: "eps"
+ TIMESTEPS: 1000
+ ZERO_TERMINAL_SNR: False
+ SCHEDULE_ARGS:
+ NAME: "scaled_linear"
+ BETA_MIN: 0.00085
+ BETA_MAX: 0.0120
+ DIFFUSION_MODEL:
+ NAME: DiffusionUNetXL
+ PRETRAINED_MODEL:
+ IN_CHANNELS: 4
+ OUT_CHANNELS: 4
+ NUM_RES_BLOCKS: 2
+ MODEL_CHANNELS: 320
+ ATTENTION_RESOLUTIONS: [4, 2]
+ DROPOUT: 0
+ CHANNEL_MULT: [1, 2, 4]
+ CONV_RESAMPLE: True
+ DIMS: 2
+ NUM_CLASSES: sequential
+ USE_CHECKPOINT: False
+ NUM_HEADS: -1
+ NUM_HEADS_CHANNELS: 64
+ USE_SCALE_SHIFT_NORM: False
+ RESBLOCK_UPDOWN: False
+ USE_NEW_ATTENTION_ORDER: True
+ USE_SPATIAL_TRANSFORMER: True
+ TRANSFORMER_DEPTH: [1, 2, 10]
+ CONTEXT_DIM: 2048
+ DISABLE_MIDDLE_SELF_ATTN: False
+ USE_LINEAR_IN_TRANSFORMER: True
+ ADM_IN_CHANNELS: 2816
+ USE_SENTENCE_EMB: False
+ USE_WORD_MAPPING: False
+ #
+ FIRST_STAGE_MODEL:
+ NAME: AutoencoderKL
+ EMBED_DIM: 4
+ IGNORE_KEYS: [ ]
+ BATCH_SIZE: 1
+ #
+ ENCODER:
+ NAME: Encoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DOUBLE_Z: True
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ #
+ DECODER:
+ NAME: Decoder
+ CH: 128
+ OUT_CH: 3
+ NUM_RES_BLOCKS: 2
+ IN_CHANNELS: 3
+ ATTN_RESOLUTIONS: [ ]
+ CH_MULT: [ 1, 2, 4, 4 ]
+ Z_CHANNELS: 4
+ DROPOUT: 0.0
+ RESAMP_WITH_CONV: True
+ GIVE_PRE_END: False
+ TANH_OUT: False
+ #
+ COND_STAGE_MODEL:
+ NAME: GeneralConditioner
+ USE_GRAD: False
+ EMBEDDERS:
+ -
+ NAME: FrozenCLIPEmbedder
+ PRETRAINED_MODEL: hf://openai/clip-vit-large-patch14
+ TOKENIZER_PATH: hf://openai/clip-vit-large-patch14
+ MAX_LENGTH: 77
+ FREEZE: True
+ LAYER: hidden
+ LAYER_IDX: 11
+ USE_FINAL_LAYER_NORM: False
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["prompt"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: FrozenOpenCLIPEmbedder2
+ ARCH: ViT-bigG-14
+ MAX_LENGTH: 77
+ FREEZE: True
+ ALWAYS_RETURN_POOLED: True
+ LEGACY: False
+ LAYER: penultimate
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["prompt"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["original_size_as_tuple"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["crop_coords_top_left"]
+ LEGACY_UCG_VALUE:
+ -
+ NAME: ConcatTimestepEmbedderND
+ OUT_DIM: 256
+ UCG_RATE: 0.0
+ INPUT_KEYS: ["target_size_as_tuple"]
+ LEGACY_UCG_VALUE:
\ No newline at end of file
diff --git a/scepter/workflow/config/tuner_model.yaml b/scepter/workflow/config/tuner_model.yaml
new file mode 100644
index 0000000..ebbe76a
--- /dev/null
+++ b/scepter/workflow/config/tuner_model.yaml
@@ -0,0 +1,561 @@
+TUNERS:
+ -
+ NAME: SD_XL1.0_Azure-Dragon
+ NAME_ZH: 青龙
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/azure_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/azure_dragon/xl_azure_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: Azure Dragon, 8K, high quality,Ultra High Detail.One of the Four Divine Creatures in Charge of Water.
+ -
+ NAME: SD_XL1.0_Gold-Dragon
+ NAME_ZH: 金龙
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/gold_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/gold_dragon/xl_gold_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: Chinese Gold Dragon in the clouds. Translucent Texture. Zbrush. Fuzzy Art. Exquisite Craftsmanship. 3D. 8K. Ultra High Detail
+ -
+ NAME: SD_XL1.0_SpringFestival-Dragon
+ NAME_ZH: 春节龙
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/spring_festival_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/spring_festival_dragon/xl_spring_festival_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: Chinese dragon. Spring Festival.Festive.Street.Lanterns.32K.High quality.expressive, dramatic, dreamlike and mysterious, Surrealism
+ -
+ NAME: SD_XL1.0_Red-Dragon
+ NAME_ZH: 红龙
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/red_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/red_dragon/xl_red_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: Traditional Red Dragon of China. Low Water Level. Studio Ghibli Style. Mural Illustration. White Background. High Detail
+ -
+ NAME: SD_XL1.0_ChinesePunk-Dragon
+ NAME_ZH: 中国朋克龙
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/chinese_punk_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/chinese_punk_dragon/xl_chinese_punk_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: uhd Image,Dragon,Chinese Dragon, Dunhuang Mural Style, Traditional Maritime Art Style
+ -
+ NAME: SD_XL1.0_Cute-Dragon
+ NAME_ZH: 喜庆龙
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/cute_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/cute_dragon/xl_kawaii_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: China Kawaii Dragon. Contest Winner. Minimalist Illustration. White Background. Flat Style. Digital Painting Style. Red. 32k uhd. Fun Comics. Fuzzy Art. Bold. Comic-Inspired Characters
+ -
+ NAME: SD_XL1.0_Dragon-Baby
+ NAME_ZH: 龙宝宝
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/baby_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/baby_dragon/xl_baby_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: Warm Colors, Soft,Chinese Dragon Baby, Felt Style,Dragon Baby, Best Quality, 3D Doll, Macaron Tones, Glittering Big Eyes, Winter,Dragon
+ -
+ NAME: SD_XL1.0_Sloppy-Dragon
+ NAME_ZH: 潦草龙
+ SOURCE: scepter
+ DESCRIPTION: None
+ BASE_MODEL: SD_XL1.0
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sloppy_dragon/
+ IMAGE_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sloppy_dragon/xl_sloppy_dragon.png
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: Messy Chinese Dragon,Cute, Wu Guanzhong, Rough
+ -
+ NAME: SD_XL1.0_Caricature
+ NAME_ZH: 夸张漫画
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/894f40ed44b37c3372e6a22b8ae577a4.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/Caricature
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Caricature
+ NAME_ZH: 夸张漫画
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/894f40ed44b37c3372e6a22b8ae577a4.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/Caricature
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Color Field Painting
+ NAME_ZH: 色域绘画
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/80e5b4075c572c04cbb4e48c37b8366b.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/ColorFieldPainting
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Color Field Painting
+ NAME_ZH: 色域绘画
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/80e5b4075c572c04cbb4e48c37b8366b.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/ColorFieldPainting
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Colored Pencil Art
+ NAME_ZH: 彩色铅笔艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9ae235d7f1a7c2a4edab52a5e9f9cbae.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/ColoredPencilArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Colored Pencil Art
+ NAME_ZH: 彩色铅笔艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/9ae235d7f1a7c2a4edab52a5e9f9cbae.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/ColoredPencilArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Dark Moody Atmosphere
+ NAME_ZH: 暗色忧郁氛围
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/3da915da2f5cedaf243e57e08163f35b.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/DarkMoodyAtmosphere
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Dark Moody Atmosphere
+ NAME_ZH: 暗色忧郁氛围
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/3da915da2f5cedaf243e57e08163f35b.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/DarkMoodyAtmosphere
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Dripping Paint Splatter Art
+ NAME_ZH: 滴漆溅画艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/69fd81f5983107acc3d334af62915851.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/DrippingPaintSplatterArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Dripping Paint Splatter Art
+ NAME_ZH: 滴漆溅画艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/69fd81f5983107acc3d334af62915851.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/DrippingPaintSplatterArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Faded Polaroid Photo
+ NAME_ZH: 褪色的宝丽来照片
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/f152edb4b3ca6248758b48115258ddfa.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/FadedPolaroidPhoto
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Faded Polaroid Photo
+ NAME_ZH: 褪色的宝丽来照片
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/f152edb4b3ca6248758b48115258ddfa.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/FadedPolaroidPhoto
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Flat 2D Art
+ NAME_ZH: 扁平2D艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/940cfd34155634cf051e1b2942cca426.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/Flat2DArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Flat 2D Art
+ NAME_ZH: 扁平2D艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/940cfd34155634cf051e1b2942cca426.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/Flat2DArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Graffiti Art
+ NAME_ZH: 涂鸦艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/57b751b11564cb22cd49ef21f2004a5f.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/GraffitiArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Graffiti Art
+ NAME_ZH: 涂鸦艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/57b751b11564cb22cd49ef21f2004a5f.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/GraffitiArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Impressionism
+ NAME_ZH: 印象主义
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/0312b673dc6858a9864d7f45f0c5c1fc.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/Impressionism
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Impressionism
+ NAME_ZH: 印象主义
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/0312b673dc6858a9864d7f45f0c5c1fc.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/Impressionism
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Logo Design
+ NAME_ZH: 标志设计
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/9aa040b0c60d289da9610c91ad9b7c7e.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/LogoDesign
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Logo Design
+ NAME_ZH: 标志设计
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/9aa040b0c60d289da9610c91ad9b7c7e.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/LogoDesign
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Pencil Sketch Drawing
+ NAME_ZH: 铅笔素描
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a9056e1eac85e5e4fe96a93917d4cce4.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/PencilSketchDrawing
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Pencil Sketch Drawing
+ NAME_ZH: 铅笔素描
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/a9056e1eac85e5e4fe96a93917d4cce4.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/PencilSketchDrawing
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Silhouette Art
+ NAME_ZH: 剪影艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/568777f447fc02510b618152726d5002.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/SilhouetteArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Silhouette Art
+ NAME_ZH: 剪影艺术
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/568777f447fc02510b618152726d5002.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/SilhouetteArt
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Steampunk 2
+ NAME_ZH: 蒸汽朋克
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/07d7b27cd73f2d43684003563511c15b.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/Steampunk2
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Steampunk 2
+ NAME_ZH: 蒸汽朋克
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/07d7b27cd73f2d43684003563511c15b.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/Steampunk2
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Sticker Designs
+ NAME_ZH: 贴纸设计
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/2d1e9867058db2c57f2fe47530de3243.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/StickerDesigns
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Sticker Designs
+ NAME_ZH: 贴纸设计
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/2d1e9867058db2c57f2fe47530de3243.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/StickerDesigns
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_Watercolor 2
+ NAME_ZH: 水彩
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8859d532ae5901cc8457d6118fb9b7da.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/Watercolor2
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_Watercolor 2
+ NAME_ZH: 水彩
+ DESCRIPTION:
+ SOURCE: diva
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/8859d532ae5901cc8457d6118fb9b7da.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/Watercolor2
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_mre-elemental-art
+ NAME_ZH: MRE元素艺术
+ DESCRIPTION:
+ SOURCE: mre
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/5895d78cf58c1ca05178991f37cc48ff.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/mre-elemental-art
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_mre-elemental-art
+ NAME_ZH: MRE元素艺术
+ DESCRIPTION:
+ SOURCE: mre
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/5895d78cf58c1ca05178991f37cc48ff.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/mre-elemental-art
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_mre-anime
+ NAME_ZH: MRE动漫
+ DESCRIPTION:
+ SOURCE: mre
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a08149bc8e50f6bc65c0010d4cd416f8.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/mre-anime
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_mre-anime
+ NAME_ZH: MRE动漫
+ DESCRIPTION:
+ SOURCE: mre
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/a08149bc8e50f6bc65c0010d4cd416f8.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/mre-anime
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_mre-comic
+ NAME_ZH: MRE漫画书
+ DESCRIPTION:
+ SOURCE: mre
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/48c65cebf1fa4284d7b8feb619412e65.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/mre-comic
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_mre-comic
+ NAME_ZH: MRE漫画书
+ DESCRIPTION:
+ SOURCE: mre
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/48c65cebf1fa4284d7b8feb619412e65.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/mre-comic
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_sai-craft clay
+ NAME_ZH: SAI手工粘土
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/8fc51113f725f27326c4398a7457cd6d.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sai-craftclay
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_sai-craft clay
+ NAME_ZH: SAI手工粘土
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/8fc51113f725f27326c4398a7457cd6d.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/sai-craftclay
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_sai-fantasy art
+ NAME_ZH: SAI幻想艺术
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a6f8d92afcd5803dfb2ebecbc92091b6.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sai-fantasyart
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_sai-fantasy art
+ NAME_ZH: SAI幻想艺术
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/a6f8d92afcd5803dfb2ebecbc92091b6.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/sai-fantasyart
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_sai-line art
+ NAME_ZH: SAI线条艺术
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/034a51b0dd34b018be8859bf45b4f7ed.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sai-lineart
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_sai-line art
+ NAME_ZH: SAI线条艺术
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/034a51b0dd34b018be8859bf45b4f7ed.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/sai-lineart
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_sai-neonpunk
+ NAME_ZH: SAI霓虹朋克
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/7e9ed25bb34008beb5f417df63c4b2fe.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sai-neonpunk
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_sai-neonpunk
+ NAME_ZH: SAI霓虹朋克
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/7e9ed25bb34008beb5f417df63c4b2fe.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/sai-neonpunk
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_sai-origami
+ NAME_ZH: SAI折纸
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/924f46a8f276011a0953d7988e90ee25.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sai-origami
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_sai-origami
+ NAME_ZH: SAI折纸
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/924f46a8f276011a0953d7988e90ee25.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/sai-origami
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD_XL1.0_sai-pixel art
+ NAME_ZH: SAI像素艺术
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD_XL1.0
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD_XL1.0/a5ab89c0960be8c1216e65c98d92ae4a.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD_XL1.0/sai-pixelart
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
+ -
+ NAME: SD1.5_sai-pixel art
+ NAME_ZH: SAI像素艺术
+ DESCRIPTION:
+ SOURCE: sai
+ BASE_MODEL: SD1.5
+ IMAGE_PATH: ms://iic/scepter@mantra_images_jpg/SD1.5/a5ab89c0960be8c1216e65c98d92ae4a.jpg
+ MODEL_PATH: ms://iic/scepter_scedit@tuners_model/SD1.5/sai-pixelart
+ TUNER_TYPE: SwiftSCE
+ PROMPT_EXAMPLE: a boy wearing green jacket
diff --git a/scepter/workflow/constant.py b/scepter/workflow/constant.py
new file mode 100644
index 0000000..7f09750
--- /dev/null
+++ b/scepter/workflow/constant.py
@@ -0,0 +1,56 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import os
+from dataclasses import dataclass
+
+WORKFLOW_PREFIX = 'custom_nodes/ComfyUI-Scepter/'
+WORKFLOW_MODEL_PREFIX = 'models/scepter/'
+WORKFLOW_CONFIG_PATH = os.path.join(WORKFLOW_PREFIX, 'config/scepter_workflow.yaml')
+MANTRA_CONFIG_PATH = os.path.join(WORKFLOW_PREFIX, 'config/mantra.yaml')
+TUNER_CONFIG_PATH = os.path.join(WORKFLOW_PREFIX, 'config/tuner_model.yaml')
+CONTROL_CONFIG_PATH = os.path.join(WORKFLOW_PREFIX, 'config/control_model.yaml')
+ANNOTATOR_CONFIG_PATH = os.path.join(WORKFLOW_PREFIX, 'config/annotator.yaml')
+
+class WorkflowConfig(object):
+ def __init__(self):
+ # pass
+ from scepter.modules.utils.file_system import FS
+ from scepter.modules.utils.config import Config
+ self.workflow_config = Config(cfg_file=WORKFLOW_CONFIG_PATH)
+ self.mantra_config = Config(cfg_file=MANTRA_CONFIG_PATH)
+ self.tuner_config = Config(cfg_file=TUNER_CONFIG_PATH)
+ self.control_config = Config(cfg_file=CONTROL_CONFIG_PATH)
+ self.annotator_config = Config(cfg_file=ANNOTATOR_CONFIG_PATH)
+
+ if 'FILE_SYSTEMS' in self.workflow_config:
+ for fs_info in self.workflow_config['FILE_SYSTEMS']:
+ FS.init_fs_client(fs_info)
+
+ self.model_info = {
+ item['NAME']: {
+ "diffusion_model": item['DIFFUSION_MODEL'],
+ "first_stage_model": item['FIRST_STAGE_MODEL'],
+ "cond_stage_model": item['COND_STAGE_MODEL'],
+ "config_file": os.path.join(WORKFLOW_PREFIX, item['CONFIG']),
+ "config": Config(cfg_file=os.path.join(WORKFLOW_PREFIX, item['CONFIG'])),
+ } for item in self.workflow_config['BASE_MODELS']
+ }
+
+ self.mantra_info = {
+ item['NAME']: item for item in self.mantra_config['MANTRAS']
+ }
+
+ self.tuner_info = {
+ item['NAME']: item for item in self.tuner_config['TUNERS']
+ }
+
+ self.control_info = {
+ item['NAME']: item for item in self.control_config['CONTROLLERS']
+ }
+
+ self.anno_info = {
+ item['TYPE']: item for item in self.annotator_config['ANNOTATORS']
+ }
+
+global WORKFLOW_CONFIG
+WORKFLOW_CONFIG = WorkflowConfig()
\ No newline at end of file
diff --git a/scepter/workflow/control_node.py b/scepter/workflow/control_node.py
new file mode 100644
index 0000000..7602767
--- /dev/null
+++ b/scepter/workflow/control_node.py
@@ -0,0 +1,110 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import os
+import numpy as np
+from PIL import Image
+import torchvision.transforms as TT
+import torch
+
+from .constant import WORKFLOW_CONFIG
+
+
+class ControlNode:
+ def __init__(self):
+ self.annotators = {}
+ self.anno_info = WORKFLOW_CONFIG.anno_info
+ for tp, anno in self.anno_info.items():
+ self.annotators[tp] = {
+ 'cfg': anno,
+ 'device': 'offline',
+ 'model': None
+ }
+ self.control_info = WORKFLOW_CONFIG.control_info
+
+ CATEGORY = '🪄 ComfyUI-Scepter'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ 'required': {
+ 'source_image': ('IMAGE', ),
+ 'control_model': (list(s().control_info.keys()), ),
+ 'control_preprocessor': (list(s().annotators.keys()), ),
+ 'crop_type': (['CenterCrop', 'NoCrop'], ),
+ },
+ 'optional': {
+ 'control_scale': ('FLOAT', {
+ 'default': 1,
+ 'min': 0,
+ 'max': 1,
+ 'step': 0.05
+ }),
+ 'output_height': ('INT', {
+ 'default': 1024,
+ 'min': 256,
+ 'max': 2048,
+ }),
+ 'output_width': ('INT', {
+ 'default': 1024,
+ 'min': 256,
+ 'max': 2048,
+ }),
+ }
+ }
+
+ OUTPUT_NODE = True
+ RETURN_TYPES = ('CONDITIONING', 'IMAGE')
+ RETURN_NAMES = ('Result', 'Control Image')
+ FUNCTION = 'execute'
+
+ def execute(self, source_image, control_model, control_preprocessor,
+ crop_type, control_scale, output_height, output_width):
+ source_image = TT.ToPILImage()(source_image.squeeze(0).permute(2, 0, 1))
+ cond_image = self.extract_condition(source_image, control_preprocessor,
+ crop_type, output_height, output_width)
+ cond_image_pil = Image.fromarray(cond_image)
+ cond_image_show = torch.from_numpy(cond_image).float().unsqueeze(0)
+ ctr_model = self.control_info[control_model]
+ out = {
+ "control_model": ctr_model,
+ "crop_type": crop_type,
+ "control_scale": control_scale,
+ "control_cond_image": cond_image_pil
+ }
+ return (out, cond_image_show, )
+
+ def extract_condition(self, source_image, control_mode, crop_type,
+ output_height, output_width):
+ annotator = self.annotators[control_mode]
+ annotator = self.load_annotator(annotator)
+
+ if crop_type == 'CenterCrop':
+ source_image = TT.Resize(max(output_height,
+ output_width))(source_image)
+ source_image = TT.CenterCrop(
+ (output_height, output_width))(source_image)
+ cond_image = annotator['model'](np.array(source_image))
+ self.annotators[control_mode] = self.unload_annotator(annotator)
+
+ if cond_image is None:
+ raise RuntimeError('Pre-process error!')
+ return cond_image
+
+ def load_annotator(self, annotator):
+ from scepter.modules.annotator.registry import ANNOTATORS
+ from scepter.modules.utils.distribute import we
+
+ if annotator['device'] == 'offline':
+ annotator['model'] = ANNOTATORS.build(annotator['cfg'])
+ annotator['device'] = 'cpu'
+ if annotator['device'] == 'cpu':
+ annotator['model'] = annotator['model'].to(we.device_id)
+ annotator['device'] = we.device_id
+ return annotator
+
+ def unload_annotator(self, annotator):
+ if not annotator['device'] == 'offline' and not annotator[
+ 'device'] == 'cpu':
+ annotator['model'] = annotator['model'].to('cpu')
+ annotator['device'] = 'cpu'
+ return annotator
\ No newline at end of file
diff --git a/scepter/workflow/mantras_node.py b/scepter/workflow/mantras_node.py
new file mode 100644
index 0000000..a3cabed
--- /dev/null
+++ b/scepter/workflow/mantras_node.py
@@ -0,0 +1,30 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import os
+from .node_utils import load_example_image
+
+from .constant import WORKFLOW_CONFIG
+
+class MantrasNode:
+ def __init__(self):
+ self.mantra_info = WORKFLOW_CONFIG.mantra_info
+
+ CATEGORY = '🪄 ComfyUI-Scepter'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ mantras_styles = list(s().mantra_info.keys())
+ return {'required': {'mantra_styles': (mantras_styles, )}}
+
+ OUTPUT_NODE = True
+ RETURN_TYPES = ('CONDITIONING', )
+ RETURN_NAMES = ('Result', )
+ FUNCTION = 'execute'
+
+ def execute(self, mantra_styles):
+ info = self.mantra_info[mantra_styles]
+ out = {
+ 'prompt_template': info['PROMPT'],
+ 'negative_prompt_template': info['NEGATIVE_PROMPT'],
+ }
+ return (out, )
diff --git a/scepter/workflow/model_node.py b/scepter/workflow/model_node.py
new file mode 100644
index 0000000..10cb295
--- /dev/null
+++ b/scepter/workflow/model_node.py
@@ -0,0 +1,179 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import copy
+import logging
+import os
+
+from .constant import WORKFLOW_CONFIG, WORKFLOW_MODEL_PREFIX
+
+class ModelNode:
+ def __init__(self):
+ from scepter.modules.utils.logger import get_logger
+ self.pipeline = {}
+ self.diff_infer = None
+ self.cfg = WORKFLOW_CONFIG.workflow_config
+ self.model_file = WORKFLOW_CONFIG.model_info
+ self.logger = get_logger('scepter', level=logging.WARNING)
+
+ CATEGORY = '🪄 ComfyUI-Scepter'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ 'required': {
+ 'model': (list(s().model_file.keys()), ),
+ "model_source": (list(s().cfg['MODEL_SOURCE']), ),
+ 'prompt': ('STRING', {
+ 'multiline': True
+ }),
+ 'negative_prompt': ('STRING', {
+ 'multiline': True
+ })
+ },
+ 'optional': {
+ 'parameters': ('CONDITIONING', ),
+ 'mantras': ('CONDITIONING', ),
+ 'tuners': ('CONDITIONING', ),
+ 'controls': ('CONDITIONING', )
+ }
+ }
+
+ OUTPUT_NODE = True
+ RETURN_TYPES = ('IMAGE', )
+ RETURN_NAMES = ('IMAGE', )
+ FUNCTION = 'execute'
+
+ def execute(self,
+ model,
+ model_source,
+ prompt,
+ negative_prompt,
+ parameters=None,
+ mantras=None,
+ tuners=None,
+ controls=None):
+ data = self.format_parameters(model, model_source, prompt, negative_prompt,
+ parameters, mantras, tuners, controls)
+ cfg = self.model_file.get(model)['config']
+ cfg = self.source_mapping(cfg, model_source)
+ self.init_infer(model, cfg)
+ output = self.diff_infer(data[0], **data[1])
+
+ x = output['images'].permute(0, 2, 3, 1)
+ output_image = x.unsqueeze(0)
+
+ return output_image
+
+ def source_mapping(self, cfg, source, type='model'):
+ def mapping(str):
+ if source == "Local":
+ str = os.path.join(WORKFLOW_MODEL_PREFIX, str.split('/', 3)[-1].replace('@', '/'))
+ elif source == "HuggingFace":
+ str = str.replace('ms://iic/', 'hf://scepter-studio/')
+ return str
+
+ if type == 'model':
+ if source == 'ModelScope':
+ return cfg
+ elif source == 'Local':
+ cfg_new = copy.deepcopy(cfg)
+ cfg_new.MODEL = cfg_new.MODEL_LOCAL
+ return cfg_new
+ elif source == 'HuggingFace':
+ cfg_new = copy.deepcopy(cfg)
+ cfg_new.MODEL = cfg_new.MODEL_HF
+ return cfg_new
+ else:
+ raise NotImplementedError(f"Unknown model source: {source}")
+ elif type in ['mantra', 'tuner', 'control']:
+ if 'MODEL_PATH' in cfg and cfg.MODEL_PATH is not None:
+ cfg.MODEL_PATH = mapping(cfg.MODEL_PATH)
+ if 'IMAGE_PATH' in cfg and cfg.IMAGE_PATH is not None:
+ cfg.IMAGE_PATH = mapping(cfg.IMAGE_PATH)
+ return cfg
+ else:
+ raise NotImplementedError(f"Unknown model source: {source}")
+
+ def init_infer(self, model_name, cfg):
+ from scepter.modules.inference.diffusion_inference import DiffusionInference
+ from scepter.modules.inference.sd3_inference import SD3Inference
+ from scepter.modules.inference.pixart_inference import PixArtInference
+ from scepter.modules.inference.flux_inference import FluxInference
+
+ if model_name.startswith('PIXART'):
+ infer_func = PixArtInference
+ elif model_name.startswith('SD3'):
+ infer_func = SD3Inference
+ elif model_name.startswith('FLUX'):
+ infer_func = FluxInference
+ else:
+ infer_func = DiffusionInference
+
+ if model_name in self.pipeline:
+ if not isinstance(self.diff_infer, infer_func):
+ self.diff_infer.dynamic_unload(name='all')
+ diff_infer = self.pipeline[model_name]
+ diff_infer.dynamic_load(name='all')
+ else:
+ diff_infer = self.pipeline[model_name]
+ else:
+ if self.diff_infer is not None:
+ self.diff_infer.dynamic_unload(name='all')
+ diff_infer = infer_func(logger=self.logger)
+ diff_infer.init_from_cfg(cfg)
+ self.pipeline[model_name] = diff_infer
+ self.diff_infer = diff_infer
+
+ def format_parameters(self,
+ model,
+ model_source,
+ prompt,
+ negative_prompt,
+ parameters,
+ mantras,
+ tuners,
+ controls):
+ input_data = {'prompt': prompt, 'negative_prompt': negative_prompt}
+ input_params = {
+ 'diffusion_model': self.model_file.get(model)['diffusion_model'],
+ 'first_stage_model': self.model_file.get(model)['first_stage_model'],
+ 'cond_stage_model': self.model_file.get(model)['cond_stage_model']
+ }
+
+ if parameters:
+ seed = parameters.get('seed', -1)
+ input_params.update({'seed': seed})
+ input_data.update(parameters)
+
+ if mantras:
+ prompt_template = mantras['prompt_template']
+ negative_prompt_template = mantras['negative_prompt_template']
+ if prompt_template != "":
+ prompt = prompt_template.replace('{prompt}', prompt)
+ if negative_prompt_template != "":
+ negative_prompt = negative_prompt + ',' + negative_prompt_template if negative_prompt != "" else negative_prompt_template
+ input_data['prompt'] = prompt
+ input_data['negative_prompt'] = negative_prompt
+ input_params.update({'mantra_state': True})
+
+ if tuners:
+ tuner_info = tuners['tuner_info']
+ tuner_info = self.source_mapping(tuner_info, model_source, type='tuner')
+ tuner_scale = tuners['tuner_scale']
+ assert model == tuner_info['BASE_MODEL'], (
+ 'The tuner model is inconsistent with the base model, '
+ 'please ensure that the selected model is consistent')
+ input_params.update({
+ 'tuner_state': True,
+ 'tuner_model': tuner_info,
+ 'tuner_scale': tuner_scale
+ })
+
+ if controls:
+ controls['control_model'] = self.source_mapping(controls['control_model'], model_source, type='control')
+ assert model == controls['control_model']['BASE_MODEL'], (
+ 'The control model is inconsistent with the base model, '
+ 'please ensure that the selected model is consistent')
+ input_params.update(controls)
+ input_params.update({'control_state': True})
+ return [input_data, input_params]
diff --git a/scepter/workflow/node_utils.py b/scepter/workflow/node_utils.py
new file mode 100644
index 0000000..b43f588
--- /dev/null
+++ b/scepter/workflow/node_utils.py
@@ -0,0 +1,17 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+from PIL import Image
+from torchvision import transforms
+
+
+def load_example_image(content, k):
+ from scepter.modules.utils.file_system import FS
+
+ image_path = FS.get_from(content[k])
+ image = Image.open(image_path).convert('RGB')
+
+ transform = transforms.Compose([transforms.ToTensor()])
+ image_tensor = transform(image)
+ example_image = image_tensor.permute(1, 2, 0).unsqueeze(0)
+
+ return example_image
diff --git a/scepter/workflow/note_node.py b/scepter/workflow/note_node.py
new file mode 100644
index 0000000..b6b1f7a
--- /dev/null
+++ b/scepter/workflow/note_node.py
@@ -0,0 +1,16 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+
+class NoteNode:
+ def __init__(self):
+ pass
+
+ CATEGORY = '🪄 ComfyUI-Scepter'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ return {'required': {'Notebook': ('STRING', {'multiline': True})}}
+
+ OUTPUT_NODE = False
+ RETURN_TYPES = ()
+ RETURN_NAMES = ()
diff --git a/scepter/workflow/parameter_node.py b/scepter/workflow/parameter_node.py
new file mode 100644
index 0000000..8892f67
--- /dev/null
+++ b/scepter/workflow/parameter_node.py
@@ -0,0 +1,64 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import os
+
+from .constant import WORKFLOW_CONFIG
+
+class ParameterNode:
+ def __init__(self):
+ self.cfg = WORKFLOW_CONFIG.workflow_config
+
+ CATEGORY = '🪄 ComfyUI-Scepter'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ 'required': {},
+ 'optional': {
+ 'sample': (s().cfg['BASE_PARAMETERS']['SAMPLER'], ),
+ 'sample_steps': ('INT', {
+ 'default': 50,
+ 'min': 1,
+ 'max': 100,
+ 'step': 1
+ }),
+ 'guide_scale': ('FLOAT', {
+ 'default': 5,
+ 'min': 0,
+ 'max': 10,
+ 'step': 0.5
+ }),
+ 'guide_rescale': ('FLOAT', {
+ 'default': 0.5,
+ 'min': 0,
+ 'max': 1,
+ 'step': 0.1
+ }),
+ 'discretization': (s().cfg['BASE_PARAMETERS']['DISCRETIZATION'], ),
+ 'output_height': (s().cfg['BASE_PARAMETERS']['OUTPUT_HEIGHT'], ),
+ 'output_width': (s().cfg['BASE_PARAMETERS']['OUTPUT_WIDTH'], ),
+ 'random_seed': ('INT', {
+ 'default': -1,
+ 'min': -1000000,
+ 'max': 1000000
+ }),
+ }
+ }
+
+ OUTPUT_NODE = True
+ RETURN_TYPES = ('CONDITIONING', )
+ RETURN_NAMES = ('Result', )
+ FUNCTION = 'execute'
+
+ def execute(self, sample, sample_steps, guide_scale, guide_rescale,
+ discretization, output_height, output_width, random_seed):
+ out = {
+ 'sample': sample,
+ 'sample_steps': sample_steps,
+ 'guide_scale': guide_scale,
+ 'guide_rescale': guide_rescale,
+ 'discretization': discretization,
+ 'target_size_as_tuple': [output_height, output_width],
+ 'seed': random_seed
+ }
+ return (out, )
diff --git a/scepter/workflow/tuner_node.py b/scepter/workflow/tuner_node.py
new file mode 100644
index 0000000..280d8e0
--- /dev/null
+++ b/scepter/workflow/tuner_node.py
@@ -0,0 +1,44 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import os
+from .node_utils import load_example_image
+
+from .constant import WORKFLOW_CONFIG
+
+class TunerNode:
+ def __init__(self):
+ self.tuner_info = WORKFLOW_CONFIG.tuner_info
+
+ CATEGORY = '🪄 ComfyUI-Scepter'
+
+ @classmethod
+ def INPUT_TYPES(s):
+ tuner_name = list(s().tuner_info.keys())
+ return {
+ 'required': {
+ 'tuner': (tuner_name, ),
+ },
+ 'optional': {
+ 'tuner_scale': (
+ 'FLOAT',
+ {
+ 'default': 1,
+ 'min': 0,
+ 'max': 1,
+ 'step': 0.05
+ },
+ ),
+ }
+ }
+
+ OUTPUT_NODE = True
+ RETURN_TYPES = ('CONDITIONING', )
+ RETURN_NAMES = ('Result', )
+ FUNCTION = 'execute'
+
+ def execute(self, tuner, tuner_scale):
+ out = {
+ "tuner_info": self.tuner_info[tuner],
+ "tuner_scale": tuner_scale
+ }
+ return (out, )
diff --git a/tests/modules/test_diffusion_inference.py b/tests/modules/test_diffusion_inference.py
index 1396f2d..eeb869e 100644
--- a/tests/modules/test_diffusion_inference.py
+++ b/tests/modules/test_diffusion_inference.py
@@ -11,6 +11,7 @@ from PIL import Image
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.inference.diffusion_inference import DiffusionInference
from scepter.modules.inference.sd3_inference import SD3Inference
+from scepter.modules.inference.flux_inference import FluxInference
from scepter.modules.inference.stylebooth_inference import StyleboothInference
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
@@ -226,7 +227,7 @@ class DiffusionInferenceTest(unittest.TestCase):
'stylebooth_test_lowpoly_cute_dog.png')
save_image(output['images'], save_path)
- # @unittest.skip('')
+ @unittest.skip('')
def test_sd3(self):
config_file = 'scepter/methods/studio/inference/dit/sd3_pro.yaml'
cfg = Config(cfg_file=config_file)
@@ -240,6 +241,19 @@ class DiffusionInferenceTest(unittest.TestCase):
save_image(output['images'], save_path)
print(save_path)
+ # @unittest.skip('')
+ def test_flux(self):
+ config_file = 'scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml'
+ cfg = Config(cfg_file=config_file)
+ diff_infer = FluxInference(logger=self.logger)
+ diff_infer.init_from_cfg(cfg)
+ output = diff_infer({
+ 'prompt': '1 girl',
+ 'seed': 2024
+ })
+ save_path = os.path.join(self.tmp_dir, 'flux_dev_1girl.png')
+ save_image(output['images'], save_path)
+ print(save_path)
if __name__ == '__main__':
unittest.main()
diff --git a/tests/tools/test_annotators.py b/tests/tools/test_annotators.py
index b7107ab..fc83c52 100644
--- a/tests/tools/test_annotators.py
+++ b/tests/tools/test_annotators.py
@@ -4,6 +4,7 @@
import os
import unittest
+import cv2
import numpy as np
from PIL import Image
from scepter.modules.annotator.registry import ANNOTATORS
@@ -245,6 +246,244 @@ class AnnotatorTest(unittest.TestCase):
Image.fromarray(save_image).save(
os.path.join(self.save_dir, f'sunflower_processor_{key}.png'))
+ @unittest.skip('')
+ def test_annotator_doodle(self):
+ doodle_dict = {
+ 'NAME': 'DoodleAnnotator', 'PROCESSOR_TYPE': 'pidinet_sketch',
+ 'PROCESSOR_CFG': [
+ {'NAME': 'PiDiAnnotator',
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/table5_pidinet.pth'},
+ {'NAME': 'SketchAnnotator',
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/sketch_simplification_gan.pth'}
+ ]
+ }
+ doodle_anno = Config(cfg_dict=doodle_dict, load=False)
+ doodle_ins = ANNOTATORS.build(doodle_anno).to(we.device_id)
+ doodle_image = doodle_ins(self.image)
+ print("doodle's shape:", doodle_image.shape)
+ Image.fromarray(doodle_image).save(
+ os.path.join(self.save_dir, 'sunflower_doodle.png'))
+
+ @unittest.skip('')
+ def test_annotator_gray(self):
+ gray_dict = {'NAME': 'GrayAnnotator'}
+ gray_anno = Config(cfg_dict=gray_dict, load=False)
+ gray_ins = ANNOTATORS.build(gray_anno).to(we.device_id)
+ gray_image = gray_ins(self.image)
+ print("gray's shape:", gray_image.shape)
+ Image.fromarray(gray_image).save(
+ os.path.join(self.save_dir, 'sunflower_gray.png'))
+
+ @unittest.skip('')
+ def test_annotator_drawing(self):
+ cont_dict = {'NAME': 'InfoDrawContourAnnotator', 'INPUT_NC': 3, 'OUTPUT_NC': 1, 'N_RESIDUAL_BLOCKS': 3,
+ 'SIGMOID': True,
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_contour_style.pth'}
+ cont_anno = Config(cfg_dict=cont_dict, load=False)
+ cont_ins = ANNOTATORS.build(cont_anno).to(we.device_id)
+ cont_image = cont_ins(self.image)
+ print("cont's shape:", cont_image.shape)
+ Image.fromarray(cont_image).save(
+ os.path.join(self.save_dir, 'sunflower_drawing_contour_style.png'))
+
+ cont_dict = {'NAME': 'InfoDrawAnimeAnnotator', 'INPUT_NC': 3, 'OUTPUT_NC': 1, 'N_RESIDUAL_BLOCKS': 3,
+ 'SIGMOID': True,
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_anime_style.pth'}
+ cont_anno = Config(cfg_dict=cont_dict, load=False)
+ cont_ins = ANNOTATORS.build(cont_anno).to(we.device_id)
+ cont_image = cont_ins(self.image)
+ print("cont's shape:", cont_image.shape)
+ Image.fromarray(cont_image).save(
+ os.path.join(self.save_dir, 'sunflower_drawing_anime_style.png'))
+
+ cont_dict = {'NAME': 'InfoDrawOpenSketchAnnotator', 'INPUT_NC': 3, 'OUTPUT_NC': 1, 'N_RESIDUAL_BLOCKS': 3,
+ 'SIGMOID': True,
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_opensketch_style.pth'}
+ cont_anno = Config(cfg_dict=cont_dict, load=False)
+ cont_ins = ANNOTATORS.build(cont_anno).to(we.device_id)
+ cont_image = cont_ins(self.image)
+ print("cont's shape:", cont_image.shape)
+ Image.fromarray(cont_image).save(
+ os.path.join(self.save_dir, 'sunflower_drawing_opensketch_style.png'))
+
+ @unittest.skip('')
+ def test_annotator_outpainting(self):
+ outpaint_dict = {'NAME': 'OutpaintingAnnotator', 'RETURN_SOURCE': False}
+ outpaint_anno = Config(cfg_dict=outpaint_dict, load=False)
+ outpaint_ins = ANNOTATORS.build(outpaint_anno).to(we.device_id)
+ outpaint_image = outpaint_ins(self.image)
+ print("outpaint's shape:", outpaint_image.shape)
+ Image.fromarray(outpaint_image).save(
+ os.path.join(self.save_dir, 'sunflower_outpaint.png'))
+
+ outpaint_dict = {'NAME': 'OutpaintingAnnotator',
+ 'RANDOM_CFG': {'DIRECTION_RANGE': ['left', 'right', 'up'], 'RATIO_RANGE': [0.2, 0.8]}}
+ outpaint_anno = Config(cfg_dict=outpaint_dict, load=False)
+ outpaint_ins = ANNOTATORS.build(outpaint_anno).to(we.device_id)
+ outpaint_image = outpaint_ins(self.image, return_mask=True)
+ print("outpaint image's shape:", outpaint_image['image'].shape)
+ print("outpaint mask's shape:", outpaint_image['mask'].shape)
+ Image.fromarray(outpaint_image['image']).save(
+ os.path.join(self.save_dir, 'sunflower_outpaint_rand_image.png'))
+ Image.fromarray(outpaint_image['mask']).save(
+ os.path.join(self.save_dir, 'sunflower_outpaint_rand_mask.png'))
+
+ @unittest.skip('')
+ def test_annotator_inpainting(self):
+ inpaint_dict = {'NAME': 'InpaintingAnnotator'}
+ inpaint_anno = Config(cfg_dict=inpaint_dict, load=False)
+ inpaint_ins = ANNOTATORS.build(inpaint_anno).to(we.device_id)
+ mask = np.zeros_like(self.image)
+ mask = cv2.rectangle(mask, (0, 0), (150, 150), (255, 255, 255), -1)
+ mask = mask[:, :, 0] # one channel format
+ inpaint_image = inpaint_ins(self.image, mask=mask, return_mask=True)
+ print("inpaint image's shape:", inpaint_image['image'].shape)
+ print("inpaint mask's shape:", inpaint_image['mask'].shape)
+ Image.fromarray(inpaint_image['image']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_image.png'))
+ Image.fromarray(inpaint_image['mask']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_mask.png'))
+
+ inpaint_dict = {'NAME': 'InpaintingAnnotator'}
+ inpaint_anno = Config(cfg_dict=inpaint_dict, load=False)
+ inpaint_ins = ANNOTATORS.build(inpaint_anno).to(we.device_id)
+ inpaint_image = inpaint_ins(self.image, return_mask=True)
+ print("inpaint image's shape:", inpaint_image['image'].shape)
+ print("inpaint mask's shape:", inpaint_image['mask'].shape)
+ Image.fromarray(inpaint_image['image']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_image_2.png'))
+ Image.fromarray(inpaint_image['mask']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_mask_2.png'))
+
+ inpaint_dict = {'NAME': 'InpaintingAnnotator',
+ 'MASK_CFG': {"irregular_proba": 0.5,
+ "irregular_kwargs": {"min_times": 4,
+ "max_times": 10,
+ "max_width": 150,
+ "max_angle": 4,
+ "max_len": 200},
+ "box_proba": 0.5,
+ "box_kwargs": {"margin": 0,
+ "bbox_min_size": 50,
+ "bbox_max_size": 150,
+ "max_times": 5,
+ "min_times": 1}
+ }
+ }
+ inpaint_anno = Config(cfg_dict=inpaint_dict, load=False)
+ inpaint_ins = ANNOTATORS.build(inpaint_anno).to(we.device_id)
+ inpaint_image = inpaint_ins(self.image, return_mask=True)
+ print("inpaint image's shape:", inpaint_image['image'].shape)
+ print("inpaint mask's shape:", inpaint_image['mask'].shape)
+ Image.fromarray(inpaint_image['image']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_image_3.png'))
+ Image.fromarray(inpaint_image['mask']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_mask_3.png'))
+
+ inpaint_image = inpaint_ins(self.image, return_mask=True, mask_color=255)
+ print("inpaint image's shape:", inpaint_image['image'].shape)
+ print("inpaint mask's shape:", inpaint_image['mask'].shape)
+ Image.fromarray(inpaint_image['image']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_image_4.png'))
+ Image.fromarray(inpaint_image['mask']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_mask_4.png'))
+
+ inpaint_image = inpaint_ins(self.image, return_mask=True, mask_color=255, return_invert=False)
+ print("inpaint image's shape:", inpaint_image['image'].shape)
+ print("inpaint mask's shape:", inpaint_image['mask'].shape)
+ Image.fromarray(inpaint_image['image']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_image_5.png'))
+ Image.fromarray(inpaint_image['mask']).save(
+ os.path.join(self.save_dir, 'sunflower_inpaint_mask_5.png'))
+
+ @unittest.skip('')
+ def test_annotator_deg(self):
+ deg_dict = {'NAME': 'DegradationAnnotator'}
+ deg_anno = Config(cfg_dict=deg_dict, load=False)
+ deg_ins = ANNOTATORS.build(deg_anno).to(we.device_id)
+ deg_image = deg_ins(self.image)
+ print("deg's shape:", deg_image.shape)
+ Image.fromarray(deg_image).save(
+ os.path.join(self.save_dir, 'sunflower_deg.png'))
+
+ deg_dict = {
+ 'NAME': 'DegradationAnnotator',
+ 'RANDOM_DEGRADATION': True,
+ 'PARAMS': {
+ 'gaussian_noise': {},
+ 'resize': {'scale': [0.4, 0.8]},
+ 'jpeg': {'jpeg_level': [25, 75]},
+ 'gaussian_blur': {'kernel_size': [7, 9, 11, 13, 15], 'sigma': [0.9, 1.8]}
+ }
+ }
+ deg_anno = Config(cfg_dict=deg_dict, load=False)
+ deg_ins = ANNOTATORS.build(deg_anno).to(we.device_id)
+ deg_image = deg_ins(self.image)
+ print("deg's shape:", deg_image.shape)
+ Image.fromarray(deg_image).save(
+ os.path.join(self.save_dir, 'sunflower_deg_2.png'))
+
+ @unittest.skip('')
+ def test_annotator_seg(self):
+ seg_dict = {
+ 'NAME': 'ESAMAnnotator',
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/efficient_sam_vits.pt',
+ 'SAVE_MODE': 'P',
+ 'GRID_SIZE': 32,
+ }
+ seg_anno = Config(cfg_dict=seg_dict, load=False)
+ seg_ins = ANNOTATORS.build(seg_anno).to(we.device_id)
+ seg_image = seg_ins(self.image)
+ print("seg's shape:", seg_image.shape)
+ Image.fromarray(seg_image).save(
+ os.path.join(self.save_dir, 'sunflower_esam_seg.png'))
+
+ seg_dict = {
+ 'NAME': 'ESAMAnnotator',
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/efficient_sam_vits.pt',
+ 'SAVE_MODE': 'P',
+ 'GRID_SIZE': 32,
+ 'USE_DOMINANT_COLOR': True,
+ 'RETURN_MASK': True
+ }
+ seg_anno = Config(cfg_dict=seg_dict, load=False)
+ seg_ins = ANNOTATORS.build(seg_anno).to(we.device_id)
+ seg_image = seg_ins(self.image)
+ print("seg image's shape:", seg_image['image'].shape)
+ Image.fromarray(seg_image['image']).save(
+ os.path.join(self.save_dir, 'sunflower_esam_seg_dominant_image.png'))
+ print("seg mask's shape:", seg_image['mask'].shape)
+ Image.fromarray(seg_image['mask']).save(
+ os.path.join(self.save_dir, 'sunflower_esam_seg_dominant_mask.png'))
+
+ @unittest.skip('')
+ def test_annotator_samdraw(self):
+ sam_dict = {
+ 'NAME': 'SAMAnnotatorDraw',
+ 'TASK_TYPE': 'input_box',
+ 'SAM_MODEL': 'vit_b',
+ 'PRETRAINED_MODEL': 'ms://iic/scepter_annotator@annotator/ckpts/sam_vit_b_01ec64.pth'
+ }
+ sam_anno = Config(cfg_dict=sam_dict, load=False)
+ sam_ins = ANNOTATORS.build(sam_anno).to(we.device_id)
+ sam_res = sam_ins(self.image, input_box=[0, 0, 200, 200], task_type='input_box', multimask_output=False)
+ Image.fromarray(sam_res['mask']).save(os.path.join(self.save_dir, f'sunflower_sam_mask.png'))
+
+ @unittest.skip('')
+ def test_annotator_lama(self):
+ lama_dict = {
+ 'NAME': 'LamaAnnotator',
+ 'PRETRAINED_MODEL': 'ms://iic/cv_fft_inpainting_lama/'
+ }
+ lama_anno = Config(cfg_dict=lama_dict, load=False)
+ lama_ins = ANNOTATORS.build(lama_anno).to(we.device_id)
+ mask = np.zeros_like(self.image)
+ mask = cv2.rectangle(mask, (0, 0), (150, 150), (255, 255, 255), -1)
+ mask = mask[:, :, 0]
+ lama_res = lama_ins(self.image, mask)
+ print("lama's shape:", lama_res.shape)
+ Image.fromarray(lama_res).save(os.path.join(self.save_dir, f'sunflower_lama_mask2.png'))
+
if __name__ == '__main__':
unittest.main()
diff --git a/tests/utils/test_fs.py b/tests/utils/test_fs.py
index 8da327e..b587a15 100644
--- a/tests/utils/test_fs.py
+++ b/tests/utils/test_fs.py
@@ -52,6 +52,23 @@ class FSTest(unittest.TestCase):
print(f'Download from {path} to {local_path}')
self.assertTrue(os.path.exists(local_path))
+ # @unittest.skip('')
+ def test_modelscope_v2(self):
+ fs_info = {'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/cache_data'}
+ config = Config(load=False, cfg_dict=fs_info)
+ FS.init_fs_client(config)
+
+ path = 'ms://AI-ModelScope/FLUX.1-dev@dev_grid.jpg'
+ with FS.get_from(path, wait_finish=True) as local_path:
+ print(f'Download from {path} to {local_path}')
+ self.assertTrue(os.path.exists(local_path))
+
+ path = 'ms://AI-ModelScope/FLUX.1-dev@tokenizer/'
+ with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path:
+ print(f'Download from {path} to {local_path}')
+ self.assertTrue(os.path.exists(local_path))
+
+
@unittest.skip('')
def test_modelscope_token(self):
fs_info = {'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/data'}