HunYuanDiT: Add 1.2 support

This commit is contained in:
City
2024-07-04 18:01:53 +02:00
parent 62b84531f4
commit 40d2a76310
3 changed files with 54 additions and 29 deletions
+26 -9
View File
@@ -1,13 +1,6 @@
"""
List of all HYDiT model types / settings
"""
sampling_settings = {
"beta_schedule" : "linear",
"linear_start" : 0.00085,
"linear_end" : 0.03,
"timesteps" : 1000,
}
from argparse import Namespace
hydit_args = Namespace(**{ # normally from argparse
"infer_mode": "torch",
@@ -30,8 +23,32 @@ hydit_conf = {
"input_size": (1024//8, 1024//8),
"args": hydit_args,
},
"sampling_settings" : sampling_settings,
"sampling_settings" : {
"beta_schedule" : "linear",
"linear_start" : 0.00085,
"linear_end" : 0.03,
"timesteps" : 1000,
},
},
"G/2-1.2": {
"unet_config": {
"depth" : 40,
"num_heads" : 16,
"patch_size" : 2,
"hidden_size" : 1408,
"mlp_ratio" : 4.3637,
"input_size": (1024//8, 1024//8),
"cond_style": False,
"cond_res" : False,
"args": hydit_args,
},
"sampling_settings" : {
"beta_schedule" : "linear",
"linear_start" : 0.00085,
"linear_end" : 0.018,
"timesteps" : 1000,
},
}
}
# these are the same as regular DiT, I think
@@ -39,6 +56,6 @@ from ..DiT.conf import dit_conf
for name in ["XL/2", "L/2", "B/2"]:
hydit_conf[name] = {
"unet_config": dit_conf[name]["unet_config"].copy(),
"sampling_settings": sampling_settings,
"sampling_settings": hydit_conf["G/2"]["sampling_settings"],
}
hydit_conf[name]["unet_config"]["args"] = hydit_args