Added support for HelloMeme AnimateDiff model loading via simple key conversion

This commit is contained in:
Jedrzej Kosinski
2024-11-12 16:13:47 -06:00
parent c088d80a98
commit f2f9c90449
+64 -16
View File
@@ -136,6 +136,11 @@ def is_animatelcm(mm_state_dict: dict[str, Tensor]) -> bool:
return False
return True
def is_hellomeme(mm_state_dict: dict[str, Tensor]) -> bool:
for key in mm_state_dict.keys():
if "pos_embed" in key:
return True
return False
def has_conv_in(mm_state_dict: dict[str, Tensor]) -> bool:
# check if conv_in.weight and .bias are present
@@ -192,6 +197,14 @@ def find_hotshot_module_num(key: str) -> Union[int, None]:
return None
_regex_hellomeme_module_num = re.compile(r'motion_modules\.(\d+)\.')
def find_hellomeme_module_num(key: str) -> Union[int, None]:
found = _regex_hellomeme_module_num.search(key)
if found:
return int(found.group(1))
return None
def has_img_encoder(mm_state_dict: dict[str, Tensor]):
for key in mm_state_dict.keys():
if key.startswith("img_encoder."):
@@ -239,6 +252,8 @@ def normalize_ad_state_dict(mm_state_dict: dict[str, Tensor], mm_name: str) -> T
raise ValueError(f"'{mm_name}' is not a valid SD1.5 nor SDXL motion module - contained {down_block_max} downblocks.")
# determine the model's format
mm_format = AnimateDiffFormat.ANIMATEDIFF
if is_hellomeme(mm_state_dict):
convert_hellomeme_state_dict(mm_state_dict)
if is_hotshotxl(mm_state_dict):
mm_format = AnimateDiffFormat.HOTSHOTXL
if is_animatelcm(mm_state_dict):
@@ -260,6 +275,7 @@ def normalize_ad_state_dict(mm_state_dict: dict[str, Tensor], mm_name: str) -> T
if mm_format == AnimateDiffFormat.FANCYVIDEO and key in FancyVideoKeys:
continue
del mm_state_dict[key]
# determine the model's version
mm_version = AnimateDiffVersion.V1
if has_mid_block(mm_state_dict):
@@ -269,26 +285,58 @@ def normalize_ad_state_dict(mm_state_dict: dict[str, Tensor], mm_name: str) -> T
info = AnimateDiffInfo(sd_type=sd_type, mm_format=mm_format, mm_version=mm_version, mm_name=mm_name)
# convert to AnimateDiff format, if needed
if mm_format == AnimateDiffFormat.HOTSHOTXL:
# HotshotXL is AD-based architecture applied to SDXL instead of SD1.5
# By renaming the keys, no code needs to be adapted at all
#
# reformat temporal_attentions:
# HSXL: temporal_attentions.#.
# AD: motion_modules.#.temporal_transformer.
# HSXL: pos_encoder.positional_encoding
# AD: pos_encoder.pe
for key in list(mm_state_dict.keys()):
module_num = find_hotshot_module_num(key)
if module_num is not None:
new_key = key.replace(f"temporal_attentions.{module_num}",
f"motion_modules.{module_num}.temporal_transformer", 1)
new_key = new_key.replace("pos_encoder.positional_encoding", "pos_encoder.pe")
mm_state_dict[new_key] = mm_state_dict[key]
del mm_state_dict[key]
convert_hotshot_state_dict(mm_state_dict)
# return adjusted mm_state_dict and info
return mm_state_dict, info
def convert_hotshot_state_dict(mm_state_dict: dict[str, Tensor]):
# HotshotXL is AD-based architecture applied to SDXL instead of SD1.5
# By renaming the keys, no code needs to be adapted at all
################################
# reformat temporal_attentions:
# HSXL: temporal_attentions.#.
# AD: motion_modules.#.temporal_transformer.
# HSXL: pos_encoder.positional_encoding
# AD: pos_encoder.pe
for key in list(mm_state_dict.keys()):
module_num = find_hotshot_module_num(key)
if module_num is not None:
new_key = key.replace(f"temporal_attentions.{module_num}",
f"motion_modules.{module_num}.temporal_transformer", 1)
new_key = new_key.replace("pos_encoder.positional_encoding", "pos_encoder.pe")
mm_state_dict[new_key] = mm_state_dict[key]
del mm_state_dict[key]
def convert_hellomeme_state_dict(mm_state_dict: dict[str, Tensor]):
# HelloMeme is AD-based architecture
for key in list(mm_state_dict.keys()):
module_num = find_hellomeme_module_num(key)
if module_num is not None:
# first, add temporal_transformer everywhere as suffix after motion_modules.#.
new_key = key.replace(f"motion_modules.{module_num}",
f"motion_modules.{module_num}.temporal_transformer")
if "pos_embed" in new_key:
new_key1 = new_key.replace("pos_embed.pe", "attention_blocks.0.pos_encoder.pe")
new_key2 = new_key.replace("pos_embed.pe", "attention_blocks.1.pos_encoder.pe")
mm_state_dict[new_key1] = mm_state_dict[key].clone()
mm_state_dict[new_key2] = mm_state_dict[key].clone()
else:
if "attn1" in new_key:
new_key = new_key.replace("attn1.", "attention_blocks.0.")
elif "attn2" in new_key:
new_key = new_key.replace("attn2.", "attention_blocks.1.")
elif "norm1" in new_key:
new_key = new_key.replace("norm1.", "norms.0.")
elif "norm2" in new_key:
new_key = new_key.replace("norm2.", "norms.1.")
elif "norm3" in new_key:
new_key = new_key.replace("norm3.", "ff_norm.")
mm_state_dict[new_key] = mm_state_dict[key]
del mm_state_dict[key]
class BlockType:
UP = "up"
DOWN = "down"