Fix UniAnimate and MultiTalk loading

This commit is contained in:
kijai
2025-08-25 14:09:08 +03:00
parent 42df277488
commit d9def84332
3 changed files with 132 additions and 235 deletions
+1 -8
View File
@@ -22,14 +22,8 @@ class MultiTalkModelLoader:
def loadmodel(self, model, base_precision=None):
from .multitalk import AudioProjModel
offload_device = mm.unet_offload_device()
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
if model_path.endswith(".gguf"):
from diffusers.models.model_loading_utils import load_gguf_checkpoint
sd = load_gguf_checkpoint(model_path)
else:
sd = load_torch_file(model_path, device=offload_device, safe_load=True)
audio_window=5
intermediate_dim=512
@@ -50,8 +44,7 @@ class MultiTalkModelLoader:
multitalk = {
"proj_model": multitalk_proj_model,
"sd": sd,
"is_gguf": model_path.endswith(".gguf"),
"model_path": model_path,
"model_type": "InfiniteTalk" if "infinite" in model.lower() else "MultiTalk",
}
+41 -38
View File
@@ -1701,20 +1701,21 @@ class WanVideoSampler:
is_5b = transformer.out_dim == 48
vae_upscale_factor = 16 if is_5b else 8
# Load weights
if transformer.patched_linear and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
if gguf_reader is not None:
if gguf_reader is not None: #handle GGUF
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
set_lora_params_gguf(transformer, patcher.patches)
transformer.patched_linear = True
elif len(patcher.patches) != 0 and transformer.patched_linear:
elif len(patcher.patches) != 0 and transformer.patched_linear: #handle patched linear layers (unmerged loras, fp8 scaled)
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
if not merge_loras and fp8_matmul:
raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported")
set_lora_params(transformer, patcher.patches)
else:
remove_lora_from_module(transformer)
remove_lora_from_module(transformer) #clear possible unmerged lora weights
transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False)
@@ -2386,16 +2387,18 @@ class WanVideoSampler:
import copy
sample_scheduler_flipped = copy.deepcopy(sample_scheduler)
#rope
# Rotary positional embeddings (RoPE)
# RoPE base freq scaling as used with CineScale
ntk_alphas = [1.0, 1.0, 1.0]
if isinstance(rope_function, dict):
ntk_alphas = rope_function["ntk_scale_f"], rope_function["ntk_scale_h"], rope_function["ntk_scale_w"]
rope_function = rope_function["rope_function"]
freqs = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
if "default" in rope_function or bidirectional_sampling:
if "default" in rope_function or bidirectional_sampling: # original RoPE
d = transformer.dim // transformer.num_heads
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
@@ -2403,7 +2406,7 @@ class WanVideoSampler:
rope_params(1024, 2 * (d // 6))
],
dim=1)
elif "comfy" in rope_function:
elif "comfy" in rope_function: # comfy's rope
transformer.rope_embedder.k = riflex_freq_index
transformer.rope_embedder.num_frames = latent_video_length
@@ -2565,37 +2568,37 @@ class WanVideoSampler:
base_params = {
'seq_len': seq_len,
'device': device,
'freqs': freqs,
't': timestep,
'current_step': idx,
'last_step': len(timesteps) - 1 == idx,
'control_lora_enabled': control_lora_enabled,
'enhance_enabled': enhance_enabled,
'camera_embed': camera_embed,
'unianim_data': unianim_data,
'fun_ref': fun_ref_input if fun_ref_image is not None else None,
'fun_camera': control_camera_input if control_camera_latents is not None else None,
'audio_proj': audio_proj if fantasytalking_embeds is not None else None,
'audio_scale': audio_scale,
"pcd_data": pcd_data_input,
"controlnet": controlnet,
"add_cond": add_cond_input,
"nag_params": text_embeds.get("nag_params", {}),
"nag_context": text_embeds.get("nag_prompt_embeds", None),
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None,
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None,
"inner_t": [shot_len] if shot_len else None,
"standin_input": standin_input,
"fantasy_portrait_input": fantasy_portrait_input,
"phantom_ref": phantom_ref,
"reverse_time": reverse_time,
"ntk_alphas": ntk_alphas,
"mtv_motion_tokens": mtv_motion_tokens if mtv_input is not None else None,
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None,
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0,
"mtv_freqs": mtv_freqs if mtv_input is not None else None,
'seq_len': seq_len, # sequence length
'device': device, # main device
'freqs': freqs, # rope freqs
't': timestep, # current timestep
'current_step': idx, # current step
'last_step': len(timesteps) - 1 == idx, # is last step
'control_lora_enabled': control_lora_enabled, # control lora toggle for patch embed selection
'enhance_enabled': enhance_enabled, # enhance-a-video toggle
'camera_embed': camera_embed, # recammaster embedding
'unianim_data': unianim_data, # unianimate input
'fun_ref': fun_ref_input if fun_ref_image is not None else None, # Fun model reference latent
'fun_camera': control_camera_input if control_camera_latents is not None else None, # Fun model camera embed
'audio_proj': audio_proj if fantasytalking_embeds is not None else None, # FantasyTalking audio projection
'audio_scale': audio_scale, # FantasyTalking audio scale
"pcd_data": pcd_data_input, # Uni3C input
"controlnet": controlnet, # TheDenk's controlnet input
"add_cond": add_cond_input, # additional conditioning input
"nag_params": text_embeds.get("nag_params", {}), # normalized attention guidance
"nag_context": text_embeds.get("nag_prompt_embeds", None), # normalized attention guidance context
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk audio input
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk reference target masks
"inner_t": [shot_len] if shot_len else None, # inner timestep for EchoShot
"standin_input": standin_input, # Stand-in reference input
"fantasy_portrait_input": fantasy_portrait_input, # Fantasy portrait input
"phantom_ref": phantom_ref, # Phantom reference input
"reverse_time": reverse_time, # Reverse RoPE toggle
"ntk_alphas": ntk_alphas, # RoPE freq scaling values
"mtv_motion_tokens": mtv_motion_tokens if mtv_input is not None else None, # MTV-Crafter motion tokens
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None, # MTV-Crafter RoPE
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, # MTV-Crafter scaling
"mtv_freqs": mtv_freqs if mtv_input is not None else None, # MTV-Crafter extra RoPE freqs
}
batch_size = 1
+90 -189
View File
@@ -741,6 +741,13 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
log.info("Using GGUF to load and assign model weights to device...")
# Prepare sd from GGUF readers
# UniAnimate embedding weight workaround
unianimate_sd = {}
for key in sd.keys():
if "dwpose_embedding" in key or "randomref_embedding_pose" in key:
unianimate_sd[key] = sd[key]
sd = {}
all_tensors = []
for r in reader:
@@ -769,6 +776,8 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
is_gguf_quant = tensor.tensor_type not in [GGMLQuantizationType.F32, GGMLQuantizationType.F16]
weights = torch.from_numpy(tensor.data.copy()).to(load_device)
sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
sd.update(unianimate_sd)
del unianimate_sd
if not getattr(transformer, "gguf_patched", False):
transformer = _replace_with_gguf_linear(
@@ -795,7 +804,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
except Exception:
block_idx = None
if "loras" in name or "dwpose" in name or "randomref" in name:
if "loras" in name:
continue
# GGUF: skip GGUFParameter params
@@ -820,7 +829,6 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
elif vace_block_idx is not None:
if vace_block_idx >= len(transformer.vace_blocks) - block_swap_args.get("vace_blocks_to_swap", 0):
load_device = offload_device
# Set tensor to device
set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=sd[name.replace("_orig_mod.", "")])
cnt += 1
@@ -869,6 +877,7 @@ def patch_stand_in_lora(transformer, lora_sd, transformer_load_device, base_dtyp
param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype))
def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
unianimate_sd = None
#spacepxl's control LoRA patch
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
@@ -885,7 +894,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
from .unianimate.nodes import update_transformer
log.info("Unianimate LoRA detected, patching model...")
patcher.model.diffusion_model = update_transformer(patcher.model.diffusion_model, lora_sd)
patcher.model.diffusion_model, unianimate_sd = update_transformer(patcher.model.diffusion_model, lora_sd)
lora_sd = standardize_lora_key_format(lora_sd)
@@ -906,7 +915,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
del lora_sd
return patcher, control_lora
return patcher, control_lora, unianimate_sd
#region Model loading
class WanVideoModelLoader:
@@ -1033,15 +1042,18 @@ class WanVideoModelLoader:
if "vace_blocks.0.after_proj.weight" in sd and not "patch_embedding.weight" in sd:
raise ValueError("You are attempting to load a VACE module as a WanVideo model, instead you should use the vace_model input and matching T2V base model")
# currently this can be VAE or MTV-Crafter weights
if extra_model is not None:
if gguf:
if not extra_model["path"].endswith(".gguf"):
raise ValueError("With GGUF main model the extra model must also be a GGUF quantized, if the main model already has extra included, you can disconnect the extra module loader")
extra_sd, extra_reader = load_gguf(extra_model["path"])
gguf_reader.append(extra_reader)
del extra_reader
else:
extra_sd = load_torch_file(extra_model["path"], device=transformer_load_device, safe_load=True)
sd.update(extra_sd)
del extra_sd
first_key = next(iter(sd))
if first_key.startswith("model.diffusion_model."):
@@ -1206,40 +1218,54 @@ class WanVideoModelLoader:
log.info("FantasyPortrait model detected, patching model...")
context_dim = fantasyportrait_model["sd"]["ip_adapter.blocks.0.cross_attn.ip_adapter_single_stream_k_proj.weight"].shape[1]
for block in transformer.blocks:
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
with init_empty_weights():
for block in transformer.blocks:
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
ip_adapter_sd = {}
for k, v in fantasyportrait_model["sd"].items():
if k.startswith("ip_adapter."):
ip_adapter_sd[k.replace("ip_adapter.", "")] = v
sd.update(ip_adapter_sd)
del ip_adapter_sd
if multitalk_model is not None:
if multitalk_model["is_gguf"] and not gguf:
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
log.info(f"{multitalk_model_type} detected, patching model...")
multitalk_model_path = multitalk_model["model_path"]
if multitalk_model_path.endswith(".gguf") and not gguf:
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
# init audio module
from .multitalk.multitalk import SingleStreamMultiAttention
from .wanvideo.modules.model import WanLayerNorm
norm_input_visual = True #dunno what this is
for block in transformer.blocks:
block.audio_cross_attn = SingleStreamMultiAttention(
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
with init_empty_weights():
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamMultiAttention(
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
qkv_bias=True,
class_range=24,
class_interval=4,
attention_mode=attention_mode,
)
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity()
log.info(f"{multitalk_model_type} detected, patching model...")
transformer.audio_proj = multitalk_model["proj_model"]
transformer.multitalk_model_type = multitalk_model_type
sd.update(multitalk_model["sd"])
extra_model_path = multitalk_model["model_path"]
if gguf:
extra_sd, extra_reader = load_gguf(extra_model_path)
gguf_reader.append(extra_reader)
del extra_reader
else:
extra_sd = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
sd.update(extra_sd)
del extra_sd
# Additional cond latents
if "add_conv_in.weight" in sd:
def zero_module(module):
@@ -1266,178 +1292,54 @@ class WanVideoModelLoader:
patcher.model.is_patched = False
scale_weights = {}
if "fp8" in quantization:
for k, v in sd.items():
if k.endswith(".scale_weight"):
scale_weights[k] = v.to(base_dtype)
if "fp8_e4m3fn" in quantization:
weight_dtype = torch.float8_e4m3fn
elif "fp8_e5m2" in quantization:
weight_dtype = torch.float8_e5m2
else:
weight_dtype = base_dtype
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
control_lora = False
if not merge_loras and control_lora:
log.warning("Control-LoRA patching is only supported with merge_loras=True")
if lora is not None:
patcher, control_lora, unianimate_sd = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras)
if unianimate_sd is not None:
log.info("Merging UniAnimate weights to the model...")
sd.update(unianimate_sd)
del unianimate_sd
if not gguf:
if "fp8" in quantization:
for k, v in sd.items():
if k.endswith(".scale_weight"):
scale_weights[k] = v
if not merge_loras:
if merge_loras and lora is not None:
if not lora_low_mem_load:
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
if control_lora:
patch_control_lora(patcher.model.diffusion_model, device)
patcher.model.is_patched = True
log.info("Merging LoRA to the model...")
patcher = apply_lora(
patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights,)
if not control_lora:
scale_weights.clear()
patcher.patches.clear()
transformer.patched_linear = False
else:
from .custom_linear import _replace_linear
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
if "fp8_e4m3fn" in quantization:
dtype = torch.float8_e4m3fn
elif "fp8_e5m2" in quantization:
dtype = torch.float8_e5m2
else:
dtype = base_dtype
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
if not lora_low_mem_load:
log.info("Using accelerate to load and assign model weights to device...")
param_count = sum(1 for _ in transformer.named_parameters())
pbar = ProgressBar(param_count)
cnt = 0
for name, param in tqdm(transformer.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}",
total=param_count,
leave=True):
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
dtype_to_use = dtype if sd[name].dtype == dtype else dtype_to_use
if "modulation" in name or "norm" in name or "bias" in name:
dtype_to_use = base_dtype
if "patch_embedding" in name:
dtype_to_use = torch.float32
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
cnt += 1
if cnt % 100 == 0:
pbar.update(100)
transformer.patched_linear = True
#for name, param in transformer.named_parameters():
# print(name, param.dtype, param.device, param.shape)
pbar.update_absolute(param_count)
pbar.update_absolute(0)
comfy_model.diffusion_model = transformer
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
patcher.model.is_patched = False
control_lora = False
if lora is not None:
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
lora_strength = l["strength"]
if isinstance(lora_strength, list):
if merge_loras:
raise ValueError("LoRA strength should be a single value when merge_loras=True")
transformer.lora_scheduling_enabled = True
if lora_strength == 0:
log.warning(f"LoRA {lora_path} has strength 0, skipping...")
continue
lora_sd = load_torch_file(lora_path, safe_load=True)
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
from .unianimate.nodes import update_transformer
log.info("Unianimate LoRA detected, patching model...")
transformer = update_transformer(transformer, lora_sd)
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
if not any('img' in k for k in sd.keys()):
lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
#spacepxl's control LoRA patch
# for key in lora_sd.keys():
# print(key)
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
log.info("Control-LoRA detected, patching model...")
if not merge_loras:
log.warning("Control-LoRA patching is only supported with merge_loras=True, setting it to True")
merge_loras = True
control_lora = True
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
old_in_dim = transformer.in_dim # 16
new_in_dim = lora_sd["diffusion_model.patch_embedding.lora_A.weight"].shape[1]
assert new_in_dim == 32
new_in = in_cls(
new_in_dim,
transformer.patch_embedding.out_channels,
transformer.patch_embedding.kernel_size,
transformer.patch_embedding.stride,
transformer.patch_embedding.padding,
).to(device=device, dtype=torch.float32)
new_in.weight.zero_()
new_in.bias.zero_()
new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight)
new_in.bias.copy_(transformer.patch_embedding.bias)
transformer.patch_embedding = new_in
transformer.expanded_patch_embedding = new_in
if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
log.info("Stand-In LoRA detected")
for block in transformer.blocks:
block.self_attn.q_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
block.self_attn.k_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
block.self_attn.v_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
for lora in [block.self_attn.q_loras, block.self_attn.k_loras, block.self_attn.v_loras]:
for param in lora.parameters():
param.requires_grad = False
for name, param in transformer.named_parameters():
if "lora" in name:
param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype))
else:
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
del lora_sd
if not gguf and merge_loras:
log.info("Patching LoRA to the model...")
patcher = apply_lora(
patcher, device, transformer_load_device,
params_to_keep=params_to_keep, dtype=dtype, base_dtype=base_dtype, state_dict=sd,
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights)
scale_weights.clear()
patcher.patches.clear()
if gguf:
#from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
log.info("Using GGUF to load and assign model weights to device...")
param_count = sum(1 for _ in transformer.named_parameters())
out_features = sd["blocks.0.self_attn.k.weight"].shape[1]
patcher.model.diffusion_model = _replace_with_gguf_linear(patcher.model.diffusion_model, base_dtype, sd, patches=patcher.patches)
pbar = ProgressBar(param_count)
cnt = 0
for name, param in tqdm(patcher.model.diffusion_model.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}",
total=param_count,
leave=True):
if "loras" in name:
continue
#print(name, param.dtype, param.device, param.shape)
if isinstance(param, GGUFParameter):
dtype_to_use = torch.uint8
elif "patch_embedding" in name:
dtype_to_use = torch.float32
else:
dtype_to_use = base_dtype
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
cnt += 1
if cnt % 100 == 0:
pbar.update(100)
#for name, param in transformer.named_parameters():
# print(name, param.dtype, param.device, param.shape)
#patcher.load(device, full_load=True)
pbar.update_absolute(param_count)
patcher.model.is_patched = True
patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False)
if "fast" in quantization:
if lora is not None and not merge_loras:
raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs")
@@ -1577,8 +1479,6 @@ class WanVideoVAELoader:
def loadmodel(self, model_name, precision, compile_args=None):
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
#with open(os.path.join(script_directory, 'configs', 'hy_vae_config.json')) as f:
# vae_config = json.load(f)
model_path = folder_paths.get_full_path("vae", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
@@ -1592,8 +1492,9 @@ class WanVideoVAELoader:
vae = WanVideoVAE38(dtype=dtype)
vae.load_state_dict(vae_sd)
del vae_sd
vae.eval()
vae.to(device = offload_device, dtype = dtype)
vae.to(device=offload_device, dtype=dtype)
if compile_args is not None:
vae.model.decoder = torch.compile(vae.model.decoder, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])